From 38e2a621bbc7b1e8578e44f937121515d5d1505a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=9F=A9=E5=98=89=E4=B9=90?= Date: Wed, 9 Sep 2026 22:12:47 +0800 Subject: [PATCH] =?UTF-8?q?refactor(ohos):=20=E6=8B=86=E5=88=86=20OHRS=20?= =?UTF-8?q?=E5=8C=85=E5=B9=B6=E6=8C=89=20socket=20=E7=B2=BE=E7=BB=86?= =?UTF-8?q?=E4=BF=9D=E6=8A=A4=20VPN=20=E6=B5=81=E9=87=8F=20(#2543)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * refactor(ohos): split facade feature and kernel crates * feat(ohos): protect transport sockets individually * fix(ohos): keep local proxy subnets off tun * fix(ohos): expose valid config enum values * refactor(ohos): finalize reusable core boundary * test(ohos): verify split package contracts * fix(port-forward): support wildcard userspace listeners Keep the existing Host listener intact while adding a DataPlane listener for force-smoltcp IPv4 wildcard rules. Keep literal loopback destinations on the local Host path instead of exporting them through an exit node. * chore(ohos): refresh split workspace lockfile * fix(socket): normalize Windows raw socket handles * refactor(socket): carry VPN protection through host bind options * refactor(socket): simplify protection defaults and TUN ingress * refactor(socket): consolidate native protection and socket creation * fix(socket): protect outbound UDP paths * fix(socket): preserve VPN routing for RPC listeners --------- Co-authored-by: FrankHan Co-authored-by: KKRainbow <443152178@qq.com> --- .github/workflows/ohos.yml | 7 + docs/socket-protection.md | 84 +++ easytier-contrib/easytier-ohrs/Cargo.lock | 41 +- easytier-contrib/easytier-ohrs/Cargo.toml | 15 +- .../easytier-ohrs/crates/README.md | 29 ++ .../crates/easytier-ohos-core/Cargo.toml | 17 + .../crates/easytier-ohos-core/src/lib.rs | 56 ++ .../easytier-ohos-core/src}/protocol.rs | 56 +- .../easytier-ohos-core/src}/routing.rs | 72 ++- .../crates/easytier-ohos-core/src/runtime.rs | 1 + .../src/runtime/state/mod.rs | 1 + .../src/runtime/state/runtime_state.rs | 30 +- .../src/socket_protection.rs | 255 +++++++++ .../crates/easytier-ohos-features/Cargo.toml | 20 + .../easytier-ohos-features/src/config.rs | 4 + .../src/config/repository/mod.rs | 0 .../src/config/services/mod.rs | 2 + .../src/config/services/schema_service.rs | 22 +- .../src/config/services/share_link_service.rs | 11 +- .../src/config/storage/config_meta.rs | 0 .../src/config/storage/mod.rs | 1 + .../src/config/types/mod.rs | 1 + .../src/config/types/stored_config.rs | 16 - .../src/config_repo.rs | 100 ++-- .../src/config_repo/field_store.rs | 0 .../src/config_repo/import_export.rs | 0 .../src/config_repo/legacy_migration.rs | 0 .../src/config_repo/validation.rs | 0 .../crates/easytier-ohos-features/src/lib.rs | 81 +++ easytier-contrib/easytier-ohrs/src/config.rs | 4 - .../easytier-ohrs/src/config/services/mod.rs | 2 - .../easytier-ohrs/src/config/storage/mod.rs | 1 - .../easytier-ohrs/src/config/types/mod.rs | 1 - .../easytier-ohrs/src/exports/config_api.rs | 4 + .../easytier-ohrs/src/exports/runtime_api.rs | 12 +- .../easytier-ohrs/src/kernel_bridge.rs | 4 +- .../src/kernel_bridge/socket_server.rs | 8 +- easytier-contrib/easytier-ohrs/src/lib.rs | 116 +++-- .../easytier-ohrs/src/napi_types.rs | 492 ++++++++++++++++++ easytier-contrib/easytier-ohrs/src/runtime.rs | 1 - .../easytier-ohrs/src/runtime/state/mod.rs | 1 - .../src/gateway/proxy/tcp_proxy_service.rs | 1 + easytier-core/src/gateway/socks5/host.rs | 5 + easytier-core/src/host/dns.rs | 4 + easytier-core/src/host/environment.rs | 3 + easytier-core/src/host/socket/factory.rs | 9 +- easytier-core/src/host/socket/listener.rs | 3 + easytier-core/src/socket/mod.rs | 6 + easytier-core/src/socket/tcp.rs | 64 ++- easytier-core/src/socket/udp/tests.rs | 10 +- .../src/socket/udp/virtual_socket.rs | 60 ++- easytier-core/src/wasi/imports.rs | 12 +- easytier-core/src/wasi/wire/options.rs | 61 ++- easytier/src/common/dns.rs | 93 ++-- easytier/src/common/upnp.rs | 38 +- easytier/src/host_runtime.rs | 35 +- easytier/src/lib.rs | 1 + easytier/src/proto/rpc/standalone.rs | 8 +- easytier/src/socket/tcp.rs | 95 ++-- easytier/src/socket/udp.rs | 46 +- easytier/src/socket_protector.rs | 235 +++++++++ easytier/src/tunnel/common.rs | 48 +- easytier/src/tunnel/websocket.rs | 3 +- 63 files changed, 2011 insertions(+), 397 deletions(-) create mode 100644 docs/socket-protection.md create mode 100644 easytier-contrib/easytier-ohrs/crates/README.md create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/Cargo.toml create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/lib.rs rename easytier-contrib/easytier-ohrs/{src/kernel_bridge => crates/easytier-ohos-core/src}/protocol.rs (63%) rename easytier-contrib/easytier-ohrs/{src/kernel_bridge => crates/easytier-ohos-core/src}/routing.rs (51%) create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime.rs create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/mod.rs rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-core}/src/runtime/state/runtime_state.rs (95%) create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/socket_protection.rs create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/Cargo.toml create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config.rs rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config/repository/mod.rs (100%) create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/mod.rs rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config/services/schema_service.rs (94%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config/services/share_link_service.rs (93%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config/storage/config_meta.rs (100%) create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/storage/mod.rs create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/mod.rs rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config/types/stored_config.rs (79%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config_repo.rs (80%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config_repo/field_store.rs (100%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config_repo/import_export.rs (100%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config_repo/legacy_migration.rs (100%) rename easytier-contrib/easytier-ohrs/{ => crates/easytier-ohos-features}/src/config_repo/validation.rs (100%) create mode 100644 easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/lib.rs delete mode 100644 easytier-contrib/easytier-ohrs/src/config.rs delete mode 100644 easytier-contrib/easytier-ohrs/src/config/services/mod.rs delete mode 100644 easytier-contrib/easytier-ohrs/src/config/storage/mod.rs delete mode 100644 easytier-contrib/easytier-ohrs/src/config/types/mod.rs create mode 100644 easytier-contrib/easytier-ohrs/src/napi_types.rs delete mode 100644 easytier-contrib/easytier-ohrs/src/runtime.rs delete mode 100644 easytier-contrib/easytier-ohrs/src/runtime/state/mod.rs create mode 100644 easytier/src/socket_protector.rs diff --git a/.github/workflows/ohos.yml b/.github/workflows/ohos.yml index a393b7bf..9a4a5db5 100644 --- a/.github/workflows/ohos.yml +++ b/.github/workflows/ohos.yml @@ -67,6 +67,13 @@ jobs: rustup component add rustfmt cargo fmt --all --manifest-path \ easytier-contrib/easytier-ohrs/Cargo.toml -- --check + cargo test --locked --manifest-path \ + easytier-contrib/easytier-ohrs/Cargo.toml \ + -p easytier-ohos-core -p easytier-ohos-features \ + --lib -- --test-threads=1 + cargo check --locked --manifest-path \ + easytier-contrib/easytier-ohrs/Cargo.toml \ + -p easytier-ohrs --tests cargo_version=$(cargo metadata --format-version 1 --no-deps \ --manifest-path easytier/Cargo.toml | jq -r '.packages[0].version') diff --git a/docs/socket-protection.md b/docs/socket-protection.md new file mode 100644 index 00000000..8ff133cc --- /dev/null +++ b/docs/socket-protection.md @@ -0,0 +1,84 @@ +# Host socket protection + +VPN bypass is a socket-creation requirement, not an operation on a socket that +core has already connected. Core and WASI guests never need an OS file descriptor. + +## Portable requests + +`TcpBindOptions::need_protect` and `UdpBindOptions::need_protect` travel through the +existing connect/bind operations. A host requiring VPN bypass must acknowledge +successful protection before connect, bind/listen, or publishing the socket. +Failure or cancellation fails creation and discards the owned socket; emitting +an event alone is not acknowledgement. Platforms without VPN bypass can treat +the requirement as a no-op. + +- TCP connect constructors request protection, including egress for proxies. +- TCP transport/hole-punch listeners request protection. Hosts retain this flag + on the listener and protect accepted children before handing them to core. +- Local ProxyNat/SOCKS/port-forward/port-lease listeners do not request it. +- UDP transport, candidates, listeners, NAT egress and STUN request protection; + HolePunchControl, local SOCKS/port-forward and port-lease sockets do not. +- Low-level TCP/UDP bind options default to protection, including deserialization + of options without `need_protect`. The named local/TUN-facing constructors set + `false` explicitly. The UDP default keeps its historical purpose label for + socket setup; only the named `hole_punch_control()` constructor opts out. +- `with_bind()` replaces the **entire** bind object. A local listener's replacement + must retain its opt-out rather than inherit the protected default. A native + adapter must honor explicit `false`, not silently change it based on `purpose`. +- DNS and source-route queries are already host-owned operations. A bypass-enabled + host must protect their underlying sockets before querying/probing, including + DNS TCP fallback, rather than letting system DNS silently bypass this contract. + +## TUN-facing ingress and port forwarding + +The extra `force_smoltcp` wildcard port-forward ingress listener has been removed. +Normal TUN-backed ingress is delivered to the existing unprotected native listener; +its accepted sockets remain unprotected so overlay replies can return through TUN. +The physical/underlay egress socket is protected independently. A separate +DataPlane listener must not mask broken host/TUN routing in this path. + +This does not remove the existing public DataPlane listener APIs, the generic +no-TUN smoltcp TCP proxy, or `force_smoltcp` itself. Those have other uses. Whether +Android subnet proxy works without forced smoltcp needs actual platform regression +testing; socket-creation unit tests alone do not establish that result. + +## Native integration and HarmonyOS + +The native `easytier` adapter implements the creation requirement using an async +`NativeSocketProtector` callback. This is a native implementation detail, not a +new portable or WASI ABI. It takes only the native handle; the creation options +select policy, without a second purpose enum. Namespace switching is confined to synchronous socket +creation and never held across the callback's await. + +The existing native `bind` builder is async and shared by TCP/UDP creation. Its +legacy direct-call default remains unprotected; portable factories explicitly +pass their core bind options (default protected). Callers must await `.call()`. +TCP listeners reuse TCP socket creation then listen, instead of duplicating the +setup. Existing legacy WebSocket direct-call policy is unchanged. + +The HarmonyOS broker wakes the already-pending request consumer with `Notify` +(no polling timer). It keeps a duplicate FD alive until ArkTS completes +`VpnConnection.protect(fd)` and returns its ACK. The waiting creation future is +woken immediately by the oneshot acknowledgement. Failure stays fail-closed; +shutdown retains dispatched FDs for late ACKs to prevent FD reuse races. +The ArkTS request shape is unchanged; its diagnostic `purpose` string is now +the generic `"socket"`. Neither ACK routing nor protection policy uses that label. + +This guarantees ordering, not a wall-clock real-time bound: OS/ArkTS scheduling +can still delay protection. Such a delay keeps the socket unconnected; it must +never allow the first SYN/query to race ahead of protection. + +## WASI option format + +The `easytier_host` import names and function signatures are unchanged. TCP +connect, TCP listen and UDP bind use option document **version 3** (previously 2): +one `u8` boolean `need_protect` is inserted immediately after the existing purpose +byte and before the optional bind-device field. All other field encodings and +purpose values are unchanged. The document version lets older hosts reject +unsupported options instead of silently ignoring protection. DNS, environment, +instance-config and data-plane layouts/versions are unchanged. + +An embedding host must update its versioned option decoder and honor the flag +inside its existing creation implementation. The external host implementation +is not in this repository: building the guest proves propagation/compatibility +of imports, not that every external host has implemented platform protection. diff --git a/easytier-contrib/easytier-ohrs/Cargo.lock b/easytier-contrib/easytier-ohrs/Cargo.lock index 178f33bf..fff93a4d 100644 --- a/easytier-contrib/easytier-ohrs/Cargo.lock +++ b/easytier-contrib/easytier-ohrs/Cargo.lock @@ -1271,27 +1271,56 @@ dependencies = [ "zstd", ] +[[package]] +name = "easytier-ohos-core" +version = "0.1.0" +dependencies = [ + "async-trait", + "easytier", + "easytier-core", + "ipnet", + "once_cell", + "serde", + "serde_json", + "tokio", + "url", +] + +[[package]] +name = "easytier-ohos-features" +version = "0.1.0" +dependencies = [ + "base64 0.22.1", + "easytier", + "flate2", + "gethostname 1.1.0", + "once_cell", + "prost-reflect", + "rusqlite", + "serde", + "serde_json", + "tracing", + "url", + "uuid", +] + [[package]] name = "easytier-ohrs" version = "0.1.0" dependencies = [ "anyhow", "async-trait", - "base64 0.22.1", "bytes", "easytier", "easytier-core", + "easytier-ohos-core", + "easytier-ohos-features", "easytier-proto", - "flate2", "futures", - "gethostname 1.1.0", - "ipnet", "napi-build-ohos", "napi-derive-ohos", "napi-ohos", "once_cell", - "prost-reflect", - "rusqlite", "serde", "serde_json", "tokio", diff --git a/easytier-contrib/easytier-ohrs/Cargo.toml b/easytier-contrib/easytier-ohrs/Cargo.toml index 9b6d2d5e..8d7758d9 100644 --- a/easytier-contrib/easytier-ohrs/Cargo.toml +++ b/easytier-contrib/easytier-ohrs/Cargo.toml @@ -1,3 +1,10 @@ +[workspace] +members = [ + "crates/easytier-ohos-features", + "crates/easytier-ohos-core", +] +resolver = "2" + [package] name = "easytier-ohrs" version = "0.1.0" @@ -9,17 +16,16 @@ crate-type=["cdylib"] [dependencies] anyhow = "1.0" async-trait = "0.1" -base64 = "0.22" bytes = "1.5" easytier-core = { path = "../../easytier-core", default-features = false } +easytier-ohos-features = { path = "crates/easytier-ohos-features" } +easytier-ohos-core = { path = "crates/easytier-ohos-core" } easytier-proto = { path = "../../easytier-proto", default-features = false, features = [ "api", "core", "json-rpc", ] } -flate2 = "1.1" futures = "0.3" -gethostname = "1.1" easytier = { path = "../../easytier" } napi-derive-ohos = "1.1" napi-ohos = { version = "1.1", default-features = false, features = [ @@ -38,11 +44,8 @@ napi-ohos = { version = "1.1", default-features = false, features = [ "web_stream", ] } once_cell = "1.21.3" -ipnet = "2.10" serde = { version = "1.0", features = ["derive"] } serde_json = "1.0.125" -prost-reflect = { version = "0.14.5", default-features = false, features = ["derive"] } -rusqlite = { version = "0.32", features = ["bundled"] } tracing-subscriber = "0.3.19" tracing-core = "0.1.33" tracing = "0.1.41" diff --git a/easytier-contrib/easytier-ohrs/crates/README.md b/easytier-contrib/easytier-ohrs/crates/README.md new file mode 100644 index 00000000..377ca2fe --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/README.md @@ -0,0 +1,29 @@ +# HarmonyOS Rust package boundaries + +`easytier-ohrs` keeps the single `.so`/HAR and N-API compatibility surface consumed by ArkTS, but its Rust implementation is split by responsibility: + +- **`easytier-ohos-core`** owns the process Tokio runtime, `NativeInstanceManager`, runtime-state projections, kernel socket protocol DTOs, TUN route aggregation, and the platform handshake used to protect individual transport sockets. Code that starts, stops, observes, or translates EasyTier runtime state belongs here. +- **`easytier-ohos-features`** owns configuration persistence and migration, SQLite metadata/field storage, schema reflection, validation, import/export, and share links. Code that remains meaningful without a running EasyTier instance belongs here. +- **`easytier-ohrs`** is the platform facade. It owns N-API exports, HarmonyOS platform logging and nearby-management adapters, and the small amount of orchestration that passes a validated feature configuration into the kernel package. + +## Dependency direction + +```text +ArkTS/HAR + | +easytier-ohrs (N-API facade) + | | + v v +easytier-ohos-core easytier-ohos-features +``` + +The facade passes owned EasyTier configuration values into the kernel package when runtime state must be projected. The kernel and feature packages do not depend on each other, so another HarmonyOS application can reuse the kernel integration without pulling in this client's SQLite repository, migrations, schema UI metadata, or share-link services. + +## Boundary rules + +1. SQLite, schema reflection, import/export, and share-link code must not enter `easytier-ohos-core`. +2. Tokio runtime ownership, instance lifecycle, TUN attachment, kernel protocol, and runtime-state conversion must not enter `easytier-ohos-features`. +3. New ArkTS exports remain in the outer `easytier-ohrs` facade so the HAR continues to expose one stable native module. +4. Cross-package values should be owned DTOs/snapshots; feature code must not receive runtime-manager handles. +5. Kernel code must not read the feature package's repository or global storage state; the facade supplies the validated runtime values it needs. +6. The split is semantic and architectural. It is not presented as a configuration-page frame-time optimization. diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/Cargo.toml b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/Cargo.toml new file mode 100644 index 00000000..30259b67 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "easytier-ohos-core" +version = "0.1.0" +edition = "2024" +description = "HarmonyOS-side EasyTier runtime, instance lifecycle and kernel interaction state" +publish = false + +[dependencies] +async-trait = "0.1" +easytier = { path = "../../../../easytier" } +easytier-core = { path = "../../../../easytier-core", default-features = false } +ipnet = "2.10" +once_cell = "1.21.3" +serde = { version = "1.0", features = ["derive"] } +serde_json = "1.0.125" +tokio = { version = "1", features = ["rt", "rt-multi-thread", "sync", "time"] } +url = "2.5" diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/lib.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/lib.rs new file mode 100644 index 00000000..e7e81181 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/lib.rs @@ -0,0 +1,56 @@ +pub mod protocol; +pub mod routing; +pub mod runtime; +pub mod socket_protection; + +use easytier::instance::factory::{NativeInstanceManager, native_instance_manager_with_runtime}; +use once_cell::sync::Lazy; +use std::sync::Arc; +use tokio::runtime::{Builder, Runtime}; + +/// The single Tokio runtime that owns HarmonyOS kernel and web-client work. +pub static ASYNC_RUNTIME: Lazy = Lazy::new(|| { + Builder::new_multi_thread() + .enable_all() + .build() + .expect("tokio runtime for easytier-ohos-core") +}); + +/// Process-wide EasyTier instance manager. Keeping it in the kernel crate prevents feature/storage +/// code from acquiring lifecycle ownership. +pub static INSTANCE_MANAGER: Lazy> = Lazy::new(|| { + Arc::new(native_instance_manager_with_runtime( + ASYNC_RUNTIME.handle().clone(), + )) +}); + +#[cfg(test)] +mod architecture_tests { + fn assert_no_napi_annotations(path: &std::path::Path) { + for entry in std::fs::read_dir(path).expect("read source directory") { + let path = entry.expect("read source entry").path(); + if path.is_dir() { + assert_no_napi_annotations(&path); + } else if path.extension().is_some_and(|extension| extension == "rs") { + let source = std::fs::read_to_string(&path).expect("read Rust source"); + let marker = ["#[", "napi"].concat(); + assert!(!source.contains(&marker), "N-API annotation in {path:?}"); + } + } + } + + #[test] + fn inner_crate_has_no_napi_registration_dependency() { + let manifest = include_str!("../Cargo.toml"); + let runtime_dependency = ["napi", "ohos"].join("-"); + let derive_dependency = ["napi", "derive", "ohos"].join("-"); + assert!(!manifest.contains(&runtime_dependency)); + assert!(!manifest.contains(&derive_dependency)); + assert!(!manifest.contains("easytier-ohos-features")); + assert_no_napi_annotations( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("src") + .as_path(), + ); + } +} diff --git a/easytier-contrib/easytier-ohrs/src/kernel_bridge/protocol.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/protocol.rs similarity index 63% rename from easytier-contrib/easytier-ohrs/src/kernel_bridge/protocol.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/protocol.rs index e6f3eea1..1b59470f 100644 --- a/easytier-contrib/easytier-ohrs/src/kernel_bridge/protocol.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/protocol.rs @@ -1,11 +1,17 @@ -use crate::config::types::stored_config::LocalSocketSyncMessage; use serde::Serialize; use std::io::{Error, ErrorKind, Write}; use std::os::unix::net::UnixStream; #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -pub(crate) struct TunRequestPayload { +pub struct LocalSocketSyncMessage { + pub message_type: String, + pub payload_json: String, +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct TunRequestPayload { pub config_id: String, pub instance_id: String, pub display_name: String, @@ -16,7 +22,7 @@ pub(crate) struct TunRequestPayload { pub need_exit_node: bool, } -pub(crate) fn send_local_socket_message( +pub fn send_local_socket_message( stream: &mut UnixStream, message_type: &str, payload_json: String, @@ -39,7 +45,7 @@ fn shrink_clients_if_sparse(clients: &mut Vec) { } } -pub(crate) fn broadcast_local_socket_message( +pub fn broadcast_local_socket_message( clients: &mut Vec, message_type: &str, payload_json: &str, @@ -57,7 +63,7 @@ pub(crate) fn broadcast_local_socket_message( delivered } -pub(crate) fn send_local_socket_json_payload_message( +pub fn send_local_socket_json_payload_message( stream: &mut UnixStream, message_type: &str, payload_json: &str, @@ -74,7 +80,7 @@ pub(crate) fn send_local_socket_json_payload_message( Ok(()) } -pub(crate) fn broadcast_local_socket_json_payload_message( +pub fn broadcast_local_socket_json_payload_message( clients: &mut Vec, message_type: &str, payload_json: &str, @@ -91,3 +97,41 @@ pub(crate) fn broadcast_local_socket_json_payload_message( *clients = active_clients; delivered } + +#[cfg(test)] +mod tests { + use super::{broadcast_local_socket_message, send_local_socket_message}; + use std::io::Read; + use std::os::unix::net::UnixStream; + + #[test] + fn local_socket_message_uses_camel_case_newline_frame() { + let (mut sender, mut receiver) = UnixStream::pair().expect("socket pair"); + send_local_socket_message(&mut sender, "runtimeState", "{\"ok\":true}".to_string()) + .expect("send frame"); + sender + .shutdown(std::net::Shutdown::Write) + .expect("shutdown"); + + let mut raw = String::new(); + receiver.read_to_string(&mut raw).expect("read frame"); + assert_eq!( + raw, + "{\"messageType\":\"runtimeState\",\"payloadJson\":\"{\\\"ok\\\":true}\"}\n" + ); + } + + #[test] + fn broadcast_removes_disconnected_clients() { + let (sender, receiver) = UnixStream::pair().expect("socket pair"); + drop(receiver); + let mut clients = vec![sender]; + + assert!(!broadcast_local_socket_message( + &mut clients, + "runtimeState", + "{}" + )); + assert!(clients.is_empty()); + } +} diff --git a/easytier-contrib/easytier-ohrs/src/kernel_bridge/routing.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/routing.rs similarity index 51% rename from easytier-contrib/easytier-ohrs/src/kernel_bridge/routing.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/routing.rs index e910d6fa..f7a29568 100644 --- a/easytier-contrib/easytier-ohrs/src/kernel_bridge/routing.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/routing.rs @@ -1,4 +1,3 @@ -use crate::config::repository::get_runtime_config_route_overrides; use crate::runtime::state::runtime_state::RuntimeInstanceState; use ipnet::IpNet; use std::collections::HashSet; @@ -54,13 +53,11 @@ fn simplify_routes(routes: Vec) -> Vec { .collect() } -pub(crate) fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec { +pub fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec { let virtual_ipv4_cidr = instance .my_node_info .as_ref() .and_then(|info| info.virtual_ipv4_cidr.clone()); - let (manual_routes, config_proxy_cidrs) = - get_runtime_config_route_overrides(&instance.config_id); let runtime_proxy_cidrs = instance .routes .iter() @@ -72,13 +69,15 @@ pub(crate) fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec Vec { +pub fn aggregate_requested_tun_routes(instances: &[RuntimeInstanceState]) -> Vec { let mut aggregated_routes = Vec::new(); let mut seen_routes = HashSet::new(); for instance in instances.iter().filter(|instance| instance.tun_required) { @@ -90,3 +89,62 @@ pub(crate) fn aggregate_requested_tun_routes(instances: &[RuntimeInstanceState]) } aggregated_routes } + +#[cfg(test)] +mod tests { + use super::{aggregate_tun_routes, simplify_routes}; + use crate::runtime::state::runtime_state::{RouteView, runtime_instance_from_config_snapshot}; + use easytier::proto::api::manage::NetworkConfig; + + #[test] + fn simplify_routes_normalizes_deduplicates_and_removes_subnets() { + let routes = simplify_routes(vec![ + "10.0.0.7".to_string(), + "10.0.0.0/24".to_string(), + "10.0.0.42/32->peer-a".to_string(), + "2001:db8::1".to_string(), + "2001:db8::/64".to_string(), + ]); + + assert_eq!(routes, vec!["10.0.0.0/24", "2001:db8::/64"]); + } + + #[test] + fn local_proxy_cidr_is_not_installed_in_tun_routes() { + let mut instance = runtime_instance_from_config_snapshot( + "routing-test".to_string(), + "test".to_string(), + NetworkConfig { + virtual_ipv4: Some("10.144.144.1".to_string()), + network_length: Some(24), + routes: vec!["172.16.0.0/16".to_string()], + proxy_cidrs: vec!["192.168.1.0/24".to_string()], + ..Default::default() + }, + true, + ); + instance.routes.push(RouteView { + peer_id: 2, + hostname: None, + ipv4: Some("10.144.144.2".to_string()), + ipv4_cidr: Some("10.144.144.2/24".to_string()), + ipv6_cidr: None, + proxy_cidrs: vec!["10.20.0.0/16".to_string()], + next_hop_peer_id: Some(2), + cost: Some(1), + path_latency: None, + udp_nat_type: None, + tcp_nat_type: None, + inst_id: None, + version: None, + is_public_server: None, + }); + + let routes = aggregate_tun_routes(&instance); + + assert!(routes.contains(&"10.144.144.0/24".to_string())); + assert!(routes.contains(&"172.16.0.0/16".to_string())); + assert!(routes.contains(&"10.20.0.0/16".to_string())); + assert!(!routes.contains(&"192.168.1.0/24".to_string())); + } +} diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime.rs new file mode 100644 index 00000000..266c62ac --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime.rs @@ -0,0 +1 @@ +pub mod state; diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/mod.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/mod.rs new file mode 100644 index 00000000..16c67eb0 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/mod.rs @@ -0,0 +1 @@ +pub mod runtime_state; diff --git a/easytier-contrib/easytier-ohrs/src/runtime/state/runtime_state.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/runtime_state.rs similarity index 95% rename from easytier-contrib/easytier-ohrs/src/runtime/state/runtime_state.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/runtime_state.rs index baff8f1b..283e045f 100644 --- a/easytier-contrib/easytier-ohrs/src/runtime/state/runtime_state.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/runtime/state/runtime_state.rs @@ -1,5 +1,4 @@ use easytier::proto::{api, common}; -use napi_derive_ohos::napi; use serde::Serialize; use std::collections::HashSet; use std::sync::Mutex; @@ -29,7 +28,6 @@ pub fn is_tun_attached(instance_id: &str) -> bool { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct PeerConnStats { pub rx_bytes: i64, pub tx_bytes: i64, @@ -40,7 +38,6 @@ pub struct PeerConnStats { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct PeerConnInfo { pub conn_id: String, pub my_peer_id: i64, @@ -61,7 +58,6 @@ pub struct PeerConnInfo { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct PeerInfo { pub peer_id: i64, pub default_conn_id: Option, @@ -71,7 +67,6 @@ pub struct PeerInfo { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct RouteView { pub peer_id: i64, pub hostname: Option, @@ -91,7 +86,6 @@ pub struct RouteView { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct MyNodeInfo { pub virtual_ipv4: Option, pub virtual_ipv4_cidr: Option, @@ -106,7 +100,6 @@ pub struct MyNodeInfo { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct RuntimeInstanceState { pub config_id: String, pub instance_id: String, @@ -121,11 +114,12 @@ pub struct RuntimeInstanceState { pub events: Vec, pub routes: Vec, pub peers: Vec, + #[serde(skip)] + pub manual_routes: Vec, } #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct TunAggregateState { pub active: bool, pub attached_instance_ids: Vec, @@ -136,7 +130,6 @@ pub struct TunAggregateState { #[derive(Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct RuntimeAggregateState { pub instances: Vec, pub tun: TunAggregateState, @@ -324,7 +317,7 @@ fn route_to_view(route: api::instance::Route) -> RouteView { } } -pub(crate) fn peer_conn_to_view(conn: api::instance::PeerConnInfo) -> PeerConnInfo { +pub fn peer_conn_to_view(conn: api::instance::PeerConnInfo) -> PeerConnInfo { let stats = conn.stats.map(|stats| PeerConnStats { rx_bytes: stats.rx_bytes as i64, tx_bytes: stats.tx_bytes as i64, @@ -399,12 +392,22 @@ fn my_node_info_to_view(info: api::manage::MyNodeInfo) -> MyNodeInfo { pub fn runtime_instance_from_running_info( config_id: String, display_name: String, - magic_dns_enabled: bool, - need_exit_node: bool, + config: Option, info: api::manage::NetworkInstanceRunningInfo, ) -> RuntimeInstanceState { let tun_attached = info.running && is_tun_attached(&config_id); let tun_required = info.running && (info.dev_name != "no_tun" || tun_attached); + let magic_dns_enabled = config + .as_ref() + .and_then(|config| config.enable_magic_dns) + .unwrap_or(false); + let need_exit_node = config + .as_ref() + .is_some_and(|config| !config.exit_nodes.is_empty()); + let manual_routes = config + .as_ref() + .map(|config| config.routes.clone()) + .unwrap_or_default(); RuntimeInstanceState { config_id: config_id.clone(), @@ -420,6 +423,7 @@ pub fn runtime_instance_from_running_info( events: info.events, routes: info.routes.into_iter().map(route_to_view).collect(), peers: info.peers.into_iter().map(peer_to_view).collect(), + manual_routes, } } @@ -434,6 +438,7 @@ pub fn runtime_instance_from_config_snapshot( running && (config.dev_name.as_deref().unwrap_or("") != "no_tun" || tun_attached); let endpoint_urls = config_endpoint_urls(&config); let public_server_url = non_empty_string(config.public_server_url.clone()); + let manual_routes = config.routes.clone(); let my_node_info = MyNodeInfo { virtual_ipv4: non_empty_string(config.virtual_ipv4.clone()), virtual_ipv4_cidr: config_virtual_ipv4_cidr(&config), @@ -460,5 +465,6 @@ pub fn runtime_instance_from_config_snapshot( events: Vec::new(), routes: configured_route_views(&endpoint_urls, public_server_url.as_deref()), peers: configured_peer_views(&endpoint_urls), + manual_routes, } } diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/socket_protection.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/socket_protection.rs new file mode 100644 index 00000000..f433bb57 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-core/src/socket_protection.rs @@ -0,0 +1,255 @@ +use std::{ + collections::{HashMap, VecDeque}, + io, + os::fd::{AsRawFd, BorrowedFd, OwnedFd}, + sync::{Arc, Mutex}, +}; + +use async_trait::async_trait; +use easytier::socket_protector::{NativeSocketProtector, set_native_socket_protector}; +use once_cell::sync::Lazy; +use tokio::sync::{Notify, oneshot}; + +const MAX_PENDING_SOCKET_PROTECTIONS: usize = 128; + +#[derive(Debug, Clone)] +pub struct SocketProtectionRequest { + pub request_id: u64, + pub socket_fd: i32, + pub purpose: String, +} + +#[derive(Default)] +struct SocketProtectionState { + enabled: bool, + next_request_id: u64, + queued: VecDeque, + pending: HashMap, +} + +struct PendingSocketProtection { + completion: oneshot::Sender>, + _socket: OwnedFd, +} + +#[derive(Default)] +pub struct SocketProtectionManager { + state: Mutex, + request_ready: Notify, +} + +pub static SOCKET_PROTECTION_MANAGER: Lazy> = + Lazy::new(|| Arc::new(SocketProtectionManager::default())); + +impl SocketProtectionManager { + fn enable(&self) { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + state.enabled = true; + } + + fn disable(&self) { + let queued = { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + state.enabled = false; + let queued_ids = state + .queued + .drain(..) + .map(|request| request.request_id) + .collect::>(); + queued_ids + .into_iter() + .filter_map(|request_id| state.pending.remove(&request_id)) + .map(|pending| pending.completion) + .collect::>() + }; + self.request_ready.notify_waiters(); + for sender in queued { + let _ = sender.send(Err(io::Error::new( + io::ErrorKind::Interrupted, + "native socket protection stopped", + ))); + } + } + + pub async fn next_request(&self) -> Option { + loop { + let notified = self.request_ready.notified(); + { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if let Some(request) = state.queued.pop_front() { + return Some(request); + } + if !state.enabled { + return None; + } + } + notified.await; + } + } + + pub fn complete_request(&self, request_id: u64, success: bool, error: Option) -> bool { + let pending = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .pending + .remove(&request_id); + let Some(pending) = pending else { + return false; + }; + let result = if success { + Ok(()) + } else { + Err(io::Error::other(error.unwrap_or_else(|| { + "native socket protection failed".to_string() + }))) + }; + pending.completion.send(result).is_ok() + } +} + +#[async_trait] +impl NativeSocketProtector for SocketProtectionManager { + async fn protect(&self, socket_handle: u64) -> io::Result<()> { + let socket_fd = i32::try_from(socket_handle).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + format!("socket handle {socket_handle} does not fit a HarmonyOS fd"), + ) + })?; + // Keep a duplicate alive across the ArkTS Promise. Socket options set + // through it affect the same socket, while cancellation cannot turn the + // request into a stale, reused descriptor. + let protected_socket = unsafe { BorrowedFd::borrow_raw(socket_fd) }.try_clone_to_owned()?; + let protected_fd = protected_socket.as_raw_fd(); + let (sender, receiver) = oneshot::channel(); + let request_id = { + let mut state = self + .state + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !state.enabled { + return Err(io::Error::new( + io::ErrorKind::NotConnected, + "native socket protection is not active", + )); + } + if state.pending.len() >= MAX_PENDING_SOCKET_PROTECTIONS { + return Err(io::Error::new( + io::ErrorKind::WouldBlock, + "too many pending native socket protection requests", + )); + } + state.next_request_id = state.next_request_id.wrapping_add(1).max(1); + let request_id = state.next_request_id; + state.queued.push_back(SocketProtectionRequest { + request_id, + socket_fd: protected_fd, + // Keep the existing ArkTS request shape without native policy labels. + purpose: "socket".to_owned(), + }); + state.pending.insert( + request_id, + PendingSocketProtection { + completion: sender, + _socket: protected_socket, + }, + ); + request_id + }; + self.request_ready.notify_one(); + receiver.await.map_err(|_| { + io::Error::new( + io::ErrorKind::Interrupted, + format!("socket protection request {request_id} was cancelled"), + ) + })? + } +} + +pub fn enable_socket_protection() -> bool { + SOCKET_PROTECTION_MANAGER.enable(); + set_native_socket_protector(Some(SOCKET_PROTECTION_MANAGER.clone())); + true +} + +pub fn disable_socket_protection() -> bool { + set_native_socket_protector(None); + SOCKET_PROTECTION_MANAGER.disable(); + true +} + +pub fn fail_socket_protection() -> bool { + // Keep the disabled manager installed so an unexpected ArkTS pump failure + // remains fail-closed for every subsequently created transport socket. + SOCKET_PROTECTION_MANAGER.disable(); + true +} + +#[cfg(test)] +mod tests { + use super::*; + use easytier::socket_protector::NativeSocketProtector; + use std::os::fd::AsRawFd; + + #[test] + fn transport_socket_waits_for_platform_completion() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + runtime.block_on(async { + let manager = Arc::new(SocketProtectionManager::default()); + manager.enable(); + let socket = std::net::UdpSocket::bind("127.0.0.1:0").unwrap(); + let socket_fd = socket.as_raw_fd(); + let task = tokio::spawn({ + let manager = manager.clone(); + async move { manager.protect(socket_fd as u64).await } + }); + let request = manager.next_request().await.unwrap(); + assert_ne!(request.socket_fd, socket_fd); + assert!(manager.complete_request(request.request_id, true, None)); + task.await.unwrap().unwrap(); + }); + } + + #[test] + fn dispatched_socket_is_retained_during_shutdown() { + let runtime = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + runtime.block_on(async { + let manager = Arc::new(SocketProtectionManager::default()); + manager.enable(); + let socket = std::net::UdpSocket::bind("127.0.0.1:0").unwrap(); + let socket_fd = socket.as_raw_fd(); + let task = tokio::spawn({ + let manager = manager.clone(); + async move { manager.protect(socket_fd as u64).await } + }); + let request = manager.next_request().await.unwrap(); + + manager.disable(); + tokio::task::yield_now().await; + assert!(!task.is_finished()); + + assert!(manager.complete_request( + request.request_id, + false, + Some("shutdown".to_string()), + )); + assert!(task.await.unwrap().is_err()); + }); + } +} diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/Cargo.toml b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/Cargo.toml new file mode 100644 index 00000000..e97b1f9f --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "easytier-ohos-features" +version = "0.1.0" +edition = "2024" +description = "HarmonyOS-side EasyTier configuration, schema, persistence and sharing features" +publish = false + +[dependencies] +base64 = "0.22" +easytier = { path = "../../../../easytier" } +flate2 = "1.1" +gethostname = "1.1" +once_cell = "1.21.3" +prost-reflect = { version = "0.14.5", default-features = false, features = ["derive"] } +rusqlite = { version = "0.32", features = ["bundled"] } +serde = { version = "1.0", features = ["derive"] } +serde_json = "1.0.125" +tracing = "0.1.41" +url = "2.5" +uuid = { version = "1.5.0", features = ["v4", "fast-rng", "macro-diagnostics", "serde"] } diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config.rs new file mode 100644 index 00000000..371dfdd9 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config.rs @@ -0,0 +1,4 @@ +pub mod repository; +pub mod services; +pub mod storage; +pub mod types; diff --git a/easytier-contrib/easytier-ohrs/src/config/repository/mod.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/repository/mod.rs similarity index 100% rename from easytier-contrib/easytier-ohrs/src/config/repository/mod.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/repository/mod.rs diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/mod.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/mod.rs new file mode 100644 index 00000000..698af0b7 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/mod.rs @@ -0,0 +1,2 @@ +pub mod schema_service; +pub mod share_link_service; diff --git a/easytier-contrib/easytier-ohrs/src/config/services/schema_service.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/schema_service.rs similarity index 94% rename from easytier-contrib/easytier-ohrs/src/config/services/schema_service.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/schema_service.rs index d1425b56..49f03f7a 100644 --- a/easytier-contrib/easytier-ohrs/src/config/services/schema_service.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/schema_service.rs @@ -1,18 +1,15 @@ use easytier::proto::ALL_DESCRIPTOR_BYTES; -use napi_derive_ohos::napi; use once_cell::sync::Lazy; use prost_reflect::{Cardinality, DescriptorPool, FieldDescriptor, Kind, MessageDescriptor}; use serde::Serialize; #[derive(Debug, Clone, Serialize)] -#[napi(object)] pub struct FieldOption { pub label: String, pub value: String, } #[derive(Debug, Clone, Serialize)] -#[napi(object)] pub struct ValidationRule { pub rule_type: String, pub arg: String, @@ -20,7 +17,6 @@ pub struct ValidationRule { } #[derive(Debug, Clone, Serialize)] -#[napi(object)] pub struct NetworkConfigSchema { pub node_kind: String, pub name: String, @@ -38,7 +34,6 @@ pub struct NetworkConfigSchema { } #[derive(Debug, Clone, Serialize)] -#[napi(object)] pub struct ConfigFieldMapping { pub field_name: String, pub field_number: i32, @@ -119,7 +114,10 @@ fn enum_options(kind: Kind) -> Vec { .values() .map(|value| FieldOption { label: value.name().to_string(), - value: value.number().to_string(), + // Protobuf JSON uses enum names rather than numeric wire values. + // Returning the number made ArkTS write `1`, while NetworkConfig + // deserialization expects a name such as `"None"`. + value: value.name().to_string(), }) .collect(), _ => Vec::new(), @@ -410,5 +408,17 @@ mod tests { .iter() .any(|option| option.label == "PublicServer") ); + + let data_compress_algo = schema + .children + .iter() + .find(|field| field.name == "data_compress_algo") + .expect("data_compress_algo field"); + let none = data_compress_algo + .enum_options + .iter() + .find(|option| option.label == "None") + .expect("compression None option"); + assert_eq!(none.value, "None"); } } diff --git a/easytier-contrib/easytier-ohrs/src/config/services/share_link_service.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/share_link_service.rs similarity index 93% rename from easytier-contrib/easytier-ohrs/src/config/services/share_link_service.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/share_link_service.rs index 33bc65bd..27206158 100644 --- a/easytier-contrib/easytier-ohrs/src/config/services/share_link_service.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/services/share_link_service.rs @@ -162,7 +162,7 @@ pub fn import_config_share_link( #[cfg(test)] mod tests { use super::*; - use crate::config_repo::{create_config_record, init_config_store}; + use crate::config::repository::{create_config_record, init_config_store}; use std::time::{SystemTime, UNIX_EPOCH}; fn test_root() -> String { @@ -178,20 +178,21 @@ mod tests { #[test] fn share_link_roundtrip_works() { + const CONFIG_ID: &str = "00000000-0000-0000-0000-000000000003"; assert!(init_config_store(test_root())); - create_config_record("cfg-share".to_string(), "share-demo".to_string()) + create_config_record(CONFIG_ID.to_string(), "share-demo".to_string()) .expect("create config"); - let link = build_config_share_link("cfg-share", None, true).expect("share link"); + let link = build_config_share_link(CONFIG_ID, None, true).expect("share link"); let payload = parse_config_share_link(&link).expect("parse link"); let config = serde_json::from_str::(&payload.config_json).expect("config json"); assert!(payload.only_start); assert_eq!(payload.display_name.as_deref(), Some("share-demo")); - assert_ne!(config.instance_id.as_deref(), Some("cfg-share")); + assert_ne!(config.instance_id.as_deref(), Some(CONFIG_ID)); let imported_id = import_config_share_link(&link, None).expect("import link"); - assert_ne!(imported_id, "cfg-share"); + assert_ne!(imported_id, CONFIG_ID); } } diff --git a/easytier-contrib/easytier-ohrs/src/config/storage/config_meta.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/storage/config_meta.rs similarity index 100% rename from easytier-contrib/easytier-ohrs/src/config/storage/config_meta.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/storage/config_meta.rs diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/storage/mod.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/storage/mod.rs new file mode 100644 index 00000000..0202377b --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/storage/mod.rs @@ -0,0 +1 @@ +pub mod config_meta; diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/mod.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/mod.rs new file mode 100644 index 00000000..91dc9890 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/mod.rs @@ -0,0 +1 @@ +pub mod stored_config; diff --git a/easytier-contrib/easytier-ohrs/src/config/types/stored_config.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/stored_config.rs similarity index 79% rename from easytier-contrib/easytier-ohrs/src/config/types/stored_config.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/stored_config.rs index 0ba42541..df95ad5f 100644 --- a/easytier-contrib/easytier-ohrs/src/config/types/stored_config.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config/types/stored_config.rs @@ -1,9 +1,7 @@ -use napi_derive_ohos::napi; use serde::Serialize; #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct StoredConfigMeta { pub config_id: String, pub display_name: String, @@ -15,7 +13,6 @@ pub struct StoredConfigMeta { #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct StoredConfigRecord { pub meta: StoredConfigMeta, pub config_json: String, @@ -23,21 +20,18 @@ pub struct StoredConfigRecord { #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct StoredConfigList { pub configs: Vec, } #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct ExportTomlResult { pub toml_text: String, } #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct SharedConfigLinkPayload { pub config_json: String, pub display_name: Option, @@ -45,15 +39,6 @@ pub struct SharedConfigLinkPayload { } #[derive(Debug, Clone, Serialize)] -#[serde(rename_all = "camelCase")] -#[napi(object)] -pub struct LocalSocketSyncMessage { - pub message_type: String, - pub payload_json: String, -} - -#[derive(Debug, Clone, Serialize)] -#[napi(object)] pub struct KeyValuePair { pub key: String, pub value: String, @@ -61,7 +46,6 @@ pub struct KeyValuePair { #[derive(Debug, Clone, Serialize)] #[serde(rename_all = "camelCase")] -#[napi(object)] pub struct SnapshotImportResult { pub ok: bool, pub error_code: String, diff --git a/easytier-contrib/easytier-ohrs/src/config_repo.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo.rs similarity index 80% rename from easytier-contrib/easytier-ohrs/src/config_repo.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo.rs index aabdfd97..5e6c74fc 100644 --- a/easytier-contrib/easytier-ohrs/src/config_repo.rs +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo.rs @@ -4,7 +4,9 @@ use crate::config::storage::config_meta::{ reset_config_meta_store, upsert_config_meta_in_tx, }; use crate::config::types::stored_config::{ExportTomlResult, StoredConfigRecord}; +use easytier::common::config::{NetworkConfigExt, TomlConfigLoader}; use easytier::proto::api::manage::NetworkConfig; +use easytier::proto::common::CompressionAlgoPb; use once_cell::sync::Lazy; use rusqlite::params; use serde_json::Value; @@ -16,16 +18,16 @@ use std::time::Instant; static CONFIG_ROOT_DIR: Mutex> = Mutex::new(None); static RUNTIME_CONFIG_SNAPSHOTS: Lazy>> = Lazy::new(|| Mutex::new(HashMap::new())); -pub(crate) const CONFIG_DIR_NAME: &str = "easytier-configs"; -pub(crate) const KERNEL_SOCKET_FILE_NAME: &str = "easytier-kernel.sock"; +pub const CONFIG_DIR_NAME: &str = "easytier-configs"; +pub const KERNEL_SOCKET_FILE_NAME: &str = "easytier-kernel.sock"; #[derive(Clone)] -pub(crate) struct RuntimeConfigSnapshot { +pub struct RuntimeConfigSnapshot { pub display_name: String, pub config: NetworkConfig, } -pub(crate) fn cache_runtime_config_snapshot( +pub fn cache_runtime_config_snapshot( config_id: String, display_name: String, config: NetworkConfig, @@ -41,42 +43,27 @@ pub(crate) fn cache_runtime_config_snapshot( } } -pub(crate) fn clear_runtime_config_snapshot(config_id: &str) { +pub fn clear_runtime_config_snapshot(config_id: &str) { if let Ok(mut guard) = RUNTIME_CONFIG_SNAPSHOTS.lock() { guard.remove(config_id); } } -pub(crate) fn get_runtime_config_snapshot(config_id: &str) -> Option { +pub fn get_runtime_config_snapshot(config_id: &str) -> Option { RUNTIME_CONFIG_SNAPSHOTS .lock() .ok() .and_then(|guard| guard.get(config_id).cloned()) } -pub(crate) fn get_runtime_config_route_overrides(config_id: &str) -> (Vec, Vec) { - RUNTIME_CONFIG_SNAPSHOTS - .lock() - .ok() - .and_then(|guard| { - guard.get(config_id).map(|snapshot| { - ( - snapshot.config.routes.clone(), - snapshot.config.proxy_cidrs.clone(), - ) - }) - }) - .unwrap_or_default() -} - -pub(crate) fn config_root_dir() -> Option { +pub fn config_root_dir() -> Option { CONFIG_ROOT_DIR .lock() .ok() .and_then(|guard| guard.as_ref().cloned()) } -pub(crate) fn kernel_socket_path() -> Option { +pub fn kernel_socket_path() -> Option { config_root_dir().map(|root| root.join(KERNEL_SOCKET_FILE_NAME)) } @@ -280,7 +267,9 @@ pub fn set_config_field_value(config_id: &str, field: &str, json_value: &str) -> } pub fn get_default_config_json() -> Option { - crate::build_default_network_config_json().ok() + let mut config = NetworkConfig::new_from_config(TomlConfigLoader::default()).ok()?; + config.data_compress_algo = Some(CompressionAlgoPb::None as i32); + serde_json::to_string(&config).ok() } pub fn create_config_record(config_id: String, display_name: String) -> Option { @@ -292,24 +281,6 @@ pub fn create_config_record(config_id: String, display_name: String) -> Option bool { - if validation::validate_config_id(config_id).is_err() { - return false; - } - let raw = match load_config_json(config_id) { - Some(raw) => raw, - None => return false, - }; - let display_name = get_config_meta(config_id) - .map(|meta| meta.display_name) - .unwrap_or_else(|| config_id.to_string()); - let started = crate::run_network_instance_from_json(&raw); - if started && let Ok(config) = serde_json::from_str::(&raw) { - cache_runtime_config_snapshot(config_id.to_string(), display_name, config); - } - started -} - pub fn list_config_meta_json() -> String { serde_json::to_string(&list_config_meta_entries().configs).unwrap_or_else(|_| "[]".to_string()) } @@ -379,23 +350,28 @@ mod tests { #[test] fn save_get_export_delete_roundtrip() { + const CONFIG_ID: &str = "00000000-0000-0000-0000-000000000001"; let root = test_root(); assert!(init_config_store(root.clone())); - let config_json = crate::build_default_network_config_json().expect("default config"); - let saved = save_config_record("cfg-1".to_string(), "test-config".to_string(), config_json) - .expect("save config"); + let config_json = get_default_config_json().expect("default config"); + let saved = save_config_record( + CONFIG_ID.to_string(), + "test-config".to_string(), + config_json, + ) + .expect("save config"); - assert_eq!(saved.meta.config_id, "cfg-1"); + assert_eq!(saved.meta.config_id, CONFIG_ID); assert_eq!(saved.meta.display_name, "test-config"); - let loaded = get_config_record("cfg-1").expect("load config"); + let loaded = get_config_record(CONFIG_ID).expect("load config"); assert_eq!(loaded.meta.display_name, "test-config"); - assert!(loaded.config_json.contains("cfg-1")); + assert!(loaded.config_json.contains(CONFIG_ID)); let legacy_json_path = PathBuf::from(&root) .join(CONFIG_DIR_NAME) - .join("cfg-1.json"); + .join(format!("{CONFIG_ID}.json")); assert!( !legacy_json_path.exists(), "config should no longer be persisted as a per-config json file" @@ -405,52 +381,54 @@ mod tests { let field_count: i64 = conn .query_row( "SELECT COUNT(*) FROM stored_config_fields WHERE config_id = ?1", - params!["cfg-1"], + params![CONFIG_ID], |row| row.get(0), ) .expect("count config fields"); + drop(conn); assert!(field_count > 0, "config fields should be stored in sqlite"); - let exported = export_config_toml("cfg-1").expect("export toml"); + let exported = export_config_toml(CONFIG_ID).expect("export toml"); assert!(exported.toml_text.contains("instance_id")); - assert!(delete_config_record("cfg-1")); - assert!(get_config_record("cfg-1").is_none()); + assert!(delete_config_record(CONFIG_ID)); + assert!(get_config_record(CONFIG_ID).is_none()); } #[test] fn set_config_field_updates_only_requested_top_level_field() { + const CONFIG_ID: &str = "00000000-0000-0000-0000-000000000002"; let root = test_root(); assert!(init_config_store(root)); - let config_json = crate::build_default_network_config_json().expect("default config"); + let config_json = get_default_config_json().expect("default config"); save_config_record( - "cfg-field".to_string(), + CONFIG_ID.to_string(), "field-config".to_string(), config_json, ) .expect("save config"); - let before_network_name = get_config_field_value("cfg-field", "network_name"); - let before_instance_id = get_config_field_value("cfg-field", "instance_id") + let before_network_name = get_config_field_value(CONFIG_ID, "network_name"); + let before_instance_id = get_config_field_value(CONFIG_ID, "instance_id") .expect("instance id field should exist"); assert!(set_config_field_value( - "cfg-field", + CONFIG_ID, "network_name", "\"changed-network\"" )); assert_eq!( - get_config_field_value("cfg-field", "network_name"), + get_config_field_value(CONFIG_ID, "network_name"), Some("\"changed-network\"".to_string()) ); assert_eq!( - get_config_field_value("cfg-field", "instance_id"), + get_config_field_value(CONFIG_ID, "instance_id"), Some(before_instance_id) ); assert_ne!( - get_config_field_value("cfg-field", "network_name"), + get_config_field_value(CONFIG_ID, "network_name"), before_network_name ); } diff --git a/easytier-contrib/easytier-ohrs/src/config_repo/field_store.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/field_store.rs similarity index 100% rename from easytier-contrib/easytier-ohrs/src/config_repo/field_store.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/field_store.rs diff --git a/easytier-contrib/easytier-ohrs/src/config_repo/import_export.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/import_export.rs similarity index 100% rename from easytier-contrib/easytier-ohrs/src/config_repo/import_export.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/import_export.rs diff --git a/easytier-contrib/easytier-ohrs/src/config_repo/legacy_migration.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/legacy_migration.rs similarity index 100% rename from easytier-contrib/easytier-ohrs/src/config_repo/legacy_migration.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/legacy_migration.rs diff --git a/easytier-contrib/easytier-ohrs/src/config_repo/validation.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/validation.rs similarity index 100% rename from easytier-contrib/easytier-ohrs/src/config_repo/validation.rs rename to easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/config_repo/validation.rs diff --git a/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/lib.rs b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/lib.rs new file mode 100644 index 00000000..4375149a --- /dev/null +++ b/easytier-contrib/easytier-ohrs/crates/easytier-ohos-features/src/lib.rs @@ -0,0 +1,81 @@ +use std::sync::OnceLock; + +#[derive(Clone, Copy)] +pub struct FeatureLogSink { + pub enabled: fn(i32) -> bool, + pub emit: fn(i32, &str, &str), +} + +static FEATURE_LOG_SINK: OnceLock = OnceLock::new(); + +/// Installs the outer HAR facade's log sink without coupling this feature crate to N-API setup. +pub fn install_log_sink(sink: FeatureLogSink) { + let _ = FEATURE_LOG_SINK.set(sink); +} + +#[doc(hidden)] +pub fn log_enabled(level: i32) -> bool { + FEATURE_LOG_SINK + .get() + .map(|sink| (sink.enabled)(level)) + .unwrap_or(true) +} + +#[doc(hidden)] +pub fn emit_log(level: i32, message: String) { + if let Some(sink) = FEATURE_LOG_SINK.get() { + (sink.emit)(level, "RustOhrs", &message); + return; + } + match level { + 5 => tracing::error!(target: "easytier_ohrs", "{message}"), + 4 => tracing::info!(target: "easytier_ohrs", "{message}"), + _ => tracing::debug!(target: "easytier_ohrs", "{message}"), + } +} + +macro_rules! ohrs_log_error { + ($($arg:tt)*) => {{ + $crate::emit_log(5, std::format!($($arg)*)); + }}; +} + +macro_rules! ohrs_log_debug { + ($($arg:tt)*) => {{ + if $crate::log_enabled(3) { + $crate::emit_log(3, std::format!($($arg)*)); + } + }}; +} + +pub mod config; + +#[cfg(test)] +mod architecture_tests { + fn assert_no_napi_annotations(path: &std::path::Path) { + for entry in std::fs::read_dir(path).expect("read source directory") { + let path = entry.expect("read source entry").path(); + if path.is_dir() { + assert_no_napi_annotations(&path); + } else if path.extension().is_some_and(|extension| extension == "rs") { + let source = std::fs::read_to_string(&path).expect("read Rust source"); + let marker = ["#[", "napi"].concat(); + assert!(!source.contains(&marker), "N-API annotation in {path:?}"); + } + } + } + + #[test] + fn inner_crate_has_no_napi_registration_dependency() { + let manifest = include_str!("../Cargo.toml"); + let runtime_dependency = ["napi", "ohos"].join("-"); + let derive_dependency = ["napi", "derive", "ohos"].join("-"); + assert!(!manifest.contains(&runtime_dependency)); + assert!(!manifest.contains(&derive_dependency)); + assert_no_napi_annotations( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("src") + .as_path(), + ); + } +} diff --git a/easytier-contrib/easytier-ohrs/src/config.rs b/easytier-contrib/easytier-ohrs/src/config.rs deleted file mode 100644 index af649e50..00000000 --- a/easytier-contrib/easytier-ohrs/src/config.rs +++ /dev/null @@ -1,4 +0,0 @@ -pub(crate) mod repository; -pub(crate) mod services; -pub(crate) mod storage; -pub(crate) mod types; diff --git a/easytier-contrib/easytier-ohrs/src/config/services/mod.rs b/easytier-contrib/easytier-ohrs/src/config/services/mod.rs deleted file mode 100644 index 88b329de..00000000 --- a/easytier-contrib/easytier-ohrs/src/config/services/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -pub(crate) mod schema_service; -pub(crate) mod share_link_service; diff --git a/easytier-contrib/easytier-ohrs/src/config/storage/mod.rs b/easytier-contrib/easytier-ohrs/src/config/storage/mod.rs deleted file mode 100644 index 765a7267..00000000 --- a/easytier-contrib/easytier-ohrs/src/config/storage/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub(crate) mod config_meta; diff --git a/easytier-contrib/easytier-ohrs/src/config/types/mod.rs b/easytier-contrib/easytier-ohrs/src/config/types/mod.rs deleted file mode 100644 index 71a56173..00000000 --- a/easytier-contrib/easytier-ohrs/src/config/types/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub(crate) mod stored_config; diff --git a/easytier-contrib/easytier-ohrs/src/exports/config_api.rs b/easytier-contrib/easytier-ohrs/src/exports/config_api.rs index 02986cb8..3ecfce33 100644 --- a/easytier-contrib/easytier-ohrs/src/exports/config_api.rs +++ b/easytier-contrib/easytier-ohrs/src/exports/config_api.rs @@ -2,6 +2,10 @@ use crate::config; use crate::config::types::stored_config::SnapshotImportResult; pub(crate) fn init_config_store(root_dir: String) -> bool { + easytier_ohos_features::install_log_sink(easytier_ohos_features::FeatureLogSink { + enabled: crate::feature_log_enabled, + emit: crate::feature_log_sink, + }); config::repository::init_config_store(root_dir) } diff --git a/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs b/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs index 86adc30d..e5e229ef 100644 --- a/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs +++ b/easytier-contrib/easytier-ohrs/src/exports/runtime_api.rs @@ -157,19 +157,10 @@ pub(crate) fn collect_runtime_state() -> RuntimeAggregateState { .as_ref() .map(|snapshot| snapshot.display_name.clone()) .unwrap_or_else(|| config_id.clone()); - let magic_dns_enabled = snapshot - .as_ref() - .and_then(|snapshot| snapshot.config.enable_magic_dns) - .unwrap_or(false); - let need_exit_node = snapshot - .as_ref() - .map(|snapshot| !snapshot.config.exit_nodes.is_empty()) - .unwrap_or(false); instances.push(runtime_instance_from_running_info( config_id, display_name, - magic_dns_enabled, - need_exit_node, + snapshot.map(|snapshot| snapshot.config), info, )); } else if let Some(snapshot) = get_runtime_config_snapshot(&config_id) { @@ -195,6 +186,7 @@ pub(crate) fn collect_runtime_state() -> RuntimeAggregateState { events: Vec::new(), routes: Vec::new(), peers: Vec::new(), + manual_routes: Vec::new(), }); } } diff --git a/easytier-contrib/easytier-ohrs/src/kernel_bridge.rs b/easytier-contrib/easytier-ohrs/src/kernel_bridge.rs index f01868b0..3af5adde 100644 --- a/easytier-contrib/easytier-ohrs/src/kernel_bridge.rs +++ b/easytier-contrib/easytier-ohrs/src/kernel_bridge.rs @@ -1,6 +1,4 @@ -mod protocol; -mod routing; mod socket_server; -pub(crate) use routing::aggregate_requested_tun_routes; +pub(crate) use easytier_ohos_core::routing::aggregate_requested_tun_routes; pub use socket_server::{start_local_socket_server, stop_local_socket_server}; 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 ffe4e35b..912fc85b 100644 --- a/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs +++ b/easytier-contrib/easytier-ohrs/src/kernel_bridge/socket_server.rs @@ -1,15 +1,15 @@ -use super::protocol::{ - TunRequestPayload, broadcast_local_socket_json_payload_message, broadcast_local_socket_message, -}; use crate::collect_runtime_state_inner; use crate::config::repository::kernel_socket_path; -use crate::kernel_bridge::routing::aggregate_tun_routes; use crate::runtime::state::runtime_state::{ PeerConnInfo as RuntimePeerConnInfo, RuntimeAggregateState, peer_conn_to_view, }; use crate::{ASYNC_RUNTIME, INSTANCE_MANAGER}; use easytier::common::global_ctx::{EventBusSubscriber, GlobalCtxEvent}; use easytier::instance::factory::subscribe_native_instance_event; +use easytier_ohos_core::protocol::{ + TunRequestPayload, broadcast_local_socket_json_payload_message, broadcast_local_socket_message, +}; +use easytier_ohos_core::routing::aggregate_tun_routes; use once_cell::sync::Lazy; use serde::Serialize; use std::collections::{HashMap, HashSet}; diff --git a/easytier-contrib/easytier-ohrs/src/lib.rs b/easytier-contrib/easytier-ohrs/src/lib.rs index 987fadc4..30bc700b 100644 --- a/easytier-contrib/easytier-ohrs/src/lib.rs +++ b/easytier-contrib/easytier-ohrs/src/lib.rs @@ -22,28 +22,14 @@ macro_rules! ohrs_log_info { }}; } -macro_rules! ohrs_log_debug { - ($($arg:tt)*) => {{ - if $crate::platform::logging::log_manager::app_log_enabled(3) { - $crate::platform::logging::log_manager::record_app_log( - 3, - "RustOhrs", - &std::format!($($arg)*), - ); - } - }}; -} - -mod config; mod exports; mod kernel_bridge; +mod napi_types; mod nearby_management; mod platform; -mod runtime; -use config::repository::{cache_runtime_config_snapshot, start_kernel_with_config_id}; +use config::repository::cache_runtime_config_snapshot; use config::services::schema_service::{ - ConfigFieldMapping, NetworkConfigSchema, get_network_config_field_mappings as build_network_config_field_mappings, get_network_config_schema as build_network_config_schema, }; @@ -53,42 +39,42 @@ use config::services::share_link_service::{ parse_config_share_link as parse_config_share_link_inner, }; 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::{ConfigLoader, TomlConfigLoader}, }; -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}; +use easytier_ohos_core::runtime; +use easytier_ohos_core::{ASYNC_RUNTIME, INSTANCE_MANAGER}; +use easytier_ohos_features::config; use kernel_bridge::{ start_local_socket_server as start_local_socket_server_inner, stop_local_socket_server as stop_local_socket_server_inner, }; use napi_derive_ohos::napi; use napi_ohos::bindgen_prelude::Uint8Array; +use napi_types::{ + ConfigFieldMapping, KeyValuePair, NetworkConfigSchema, SharedConfigLinkPayload, + SnapshotImportResult, SocketProtectionRequest, +}; use runtime::state::runtime_state::{RuntimeAggregateState, RuntimeInstanceState}; use std::collections::{HashMap, HashSet}; use std::format; use std::sync::{Arc, Mutex}; -use tokio::runtime::{Builder, Runtime}; use uuid::Uuid; -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(), - )) - }); +pub(crate) fn feature_log_sink(level: i32, target: &str, message: &str) { + platform::logging::log_manager::record_app_log(level, target, message); +} + +pub(crate) fn feature_log_enabled(level: i32) -> bool { + platform::logging::log_manager::app_log_enabled(level) +} + static WEB_CLIENTS: once_cell::sync::Lazy>> = once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new())); const PRO_CONFIG_SERVER_CLIENT_ID: &str = "__easytier_pro_config_server_client__"; @@ -668,12 +654,6 @@ fn resolve_instance_id_inner(instance_name: &str) -> Option { resolve_instance_id_from_state(&collect_runtime_state_inner(), instance_name) } -pub(crate) fn build_default_network_config_json() -> Result { - let config = NetworkConfig::new_from_config(TomlConfigLoader::default()) - .map_err(|e| format!("default_network_config failed {}", e))?; - serde_json::to_string(&config).map_err(|e| format!("default_network_config failed {}", e)) -} - fn convert_toml_to_network_config_inner(toml_text: &str) -> Result { let config = NetworkConfig::new_from_config( TomlConfigLoader::new_from_str(toml_text).map_err(|e| e.to_string())?, @@ -750,6 +730,18 @@ pub(crate) fn run_network_instance_from_json(cfg_json: &str) -> bool { } } +fn start_kernel_with_config_id(config_id: &str) -> bool { + let Some(raw) = config::repository::load_config_json(config_id) else { + return false; + }; + let display_name = get_config_display_name(config_id).unwrap_or_else(|| config_id.to_string()); + let started = run_network_instance_from_json(&raw); + if started && let Ok(config) = serde_json::from_str::(&raw) { + cache_runtime_config_snapshot(config_id.to_string(), display_name, config); + } + started +} + fn parse_instance_uuid(config_id: &str) -> Option { match Uuid::parse_str(config_id) { Ok(uuid) => Some(uuid), @@ -847,7 +839,7 @@ pub fn import_config_store_snapshot(source_path: String) -> bool { #[napi] pub fn import_config_store_snapshot_with_result(source_path: String) -> SnapshotImportResult { - exports::config_api::import_config_store_snapshot_with_result(source_path) + exports::config_api::import_config_store_snapshot_with_result(source_path).into() } #[napi] @@ -1056,6 +1048,9 @@ pub async fn call_nearby_management_json_rpc( #[napi] pub fn collect_network_infos() -> Vec { exports::runtime_api::collect_network_infos() + .into_iter() + .map(Into::into) + .collect() } #[napi] @@ -1063,14 +1058,53 @@ pub fn set_tun_fd(config_id: String, fd: i32) -> bool { exports::runtime_api::set_tun_fd(config_id, fd, parse_instance_uuid) } +#[napi] +pub fn enable_socket_protection() -> bool { + easytier_ohos_core::socket_protection::enable_socket_protection() +} + +#[napi] +pub async fn next_socket_protection_request() -> Option { + easytier_ohos_core::socket_protection::SOCKET_PROTECTION_MANAGER + .next_request() + .await + .map(Into::into) +} + +#[napi] +pub fn complete_socket_protection( + request_id: String, + success: bool, + error: Option, +) -> bool { + let Ok(request_id) = request_id.parse::() else { + return false; + }; + easytier_ohos_core::socket_protection::SOCKET_PROTECTION_MANAGER + .complete_request(request_id, success, error) +} + +#[napi] +pub fn disable_socket_protection() -> bool { + easytier_ohos_core::socket_protection::disable_socket_protection() +} + +#[napi] +pub fn fail_socket_protection() -> bool { + easytier_ohos_core::socket_protection::fail_socket_protection() +} + #[napi] pub fn get_network_config_schema() -> NetworkConfigSchema { - build_network_config_schema() + build_network_config_schema().into() } #[napi] pub fn get_network_config_field_mappings() -> Vec { build_network_config_field_mappings() + .into_iter() + .map(Into::into) + .collect() } #[cfg(test)] @@ -1118,6 +1152,7 @@ mod tests { events: vec![], routes: vec![], peers: vec![], + manual_routes: vec![], }, RuntimeInstanceState { config_id: "ec7b6a3c-aeae-4c0e-844e-f7ec2dbdc2ce".to_string(), @@ -1133,6 +1168,7 @@ mod tests { events: vec![], routes: vec![], peers: vec![], + manual_routes: vec![], }, ], tun: runtime::state::runtime_state::TunAggregateState { @@ -1199,7 +1235,7 @@ pub fn build_config_share_link(config_id: String, only_start: Option) -> O #[napi] pub fn parse_config_share_link(share_link: String) -> Option { - parse_config_share_link_inner(&share_link) + parse_config_share_link_inner(&share_link).map(Into::into) } #[napi] diff --git a/easytier-contrib/easytier-ohrs/src/napi_types.rs b/easytier-contrib/easytier-ohrs/src/napi_types.rs new file mode 100644 index 00000000..332c3ed6 --- /dev/null +++ b/easytier-contrib/easytier-ohrs/src/napi_types.rs @@ -0,0 +1,492 @@ +#![allow(dead_code)] + +use easytier_ohos_core::runtime::state::runtime_state as kernel_types; +use easytier_ohos_features::config::services::schema_service as feature_schema; +use easytier_ohos_features::config::types::stored_config as feature_types; +use napi_derive_ohos::napi; +use serde::Serialize; + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct SocketProtectionRequest { + pub request_id: String, + pub socket_fd: i32, + pub purpose: String, +} + +impl From + for SocketProtectionRequest +{ + fn from(value: easytier_ohos_core::socket_protection::SocketProtectionRequest) -> Self { + Self { + request_id: value.request_id.to_string(), + socket_fd: value.socket_fd, + purpose: value.purpose, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct StoredConfigMeta { + pub config_id: String, + pub display_name: String, + pub created_at: String, + pub updated_at: String, + pub favorite: bool, + pub temporary: bool, +} + +impl From for StoredConfigMeta { + fn from(value: feature_types::StoredConfigMeta) -> Self { + Self { + config_id: value.config_id, + display_name: value.display_name, + created_at: value.created_at, + updated_at: value.updated_at, + favorite: value.favorite, + temporary: value.temporary, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct StoredConfigRecord { + pub meta: StoredConfigMeta, + pub config_json: String, +} + +impl From for StoredConfigRecord { + fn from(value: feature_types::StoredConfigRecord) -> Self { + Self { + meta: value.meta.into(), + config_json: value.config_json, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct StoredConfigList { + pub configs: Vec, +} + +impl From for StoredConfigList { + fn from(value: feature_types::StoredConfigList) -> Self { + Self { + configs: value.configs.into_iter().map(Into::into).collect(), + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct ExportTomlResult { + pub toml_text: String, +} + +impl From for ExportTomlResult { + fn from(value: feature_types::ExportTomlResult) -> Self { + Self { + toml_text: value.toml_text, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct SharedConfigLinkPayload { + pub config_json: String, + pub display_name: Option, + pub only_start: bool, +} + +impl From for SharedConfigLinkPayload { + fn from(value: feature_types::SharedConfigLinkPayload) -> Self { + Self { + config_json: value.config_json, + display_name: value.display_name, + only_start: value.only_start, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct LocalSocketSyncMessage { + pub message_type: String, + pub payload_json: String, +} + +#[derive(Debug, Clone, Serialize)] +#[napi(object)] +pub struct KeyValuePair { + pub key: String, + pub value: String, +} + +impl From for KeyValuePair { + fn from(value: feature_types::KeyValuePair) -> Self { + Self { + key: value.key, + value: value.value, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct SnapshotImportResult { + pub ok: bool, + pub error_code: String, + pub error_message: String, + pub snapshot_invalid: bool, +} + +impl From for SnapshotImportResult { + fn from(value: feature_types::SnapshotImportResult) -> Self { + Self { + ok: value.ok, + error_code: value.error_code, + error_message: value.error_message, + snapshot_invalid: value.snapshot_invalid, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[napi(object)] +pub struct FieldOption { + pub label: String, + pub value: String, +} + +impl From for FieldOption { + fn from(value: feature_schema::FieldOption) -> Self { + Self { + label: value.label, + value: value.value, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[napi(object)] +pub struct ValidationRule { + pub rule_type: String, + pub arg: String, + pub message: String, +} + +impl From for ValidationRule { + fn from(value: feature_schema::ValidationRule) -> Self { + Self { + rule_type: value.rule_type, + arg: value.arg, + message: value.message, + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[napi(object)] +pub struct NetworkConfigSchema { + pub node_kind: String, + pub name: String, + pub field_number: i32, + pub type_name: Option, + pub semantic_type: Option, + pub value_kind: String, + pub is_list: bool, + pub required: bool, + pub default_value_text: Option, + pub enum_options: Vec, + pub validations: Vec, + pub children: Vec, + pub definitions: Vec, +} + +impl From for NetworkConfigSchema { + fn from(value: feature_schema::NetworkConfigSchema) -> Self { + Self { + node_kind: value.node_kind, + name: value.name, + field_number: value.field_number, + type_name: value.type_name, + semantic_type: value.semantic_type, + value_kind: value.value_kind, + is_list: value.is_list, + required: value.required, + default_value_text: value.default_value_text, + enum_options: value.enum_options.into_iter().map(Into::into).collect(), + validations: value.validations.into_iter().map(Into::into).collect(), + children: value.children.into_iter().map(Into::into).collect(), + definitions: value.definitions.into_iter().map(Into::into).collect(), + } + } +} + +#[derive(Debug, Clone, Serialize)] +#[napi(object)] +pub struct ConfigFieldMapping { + pub field_name: String, + pub field_number: i32, +} + +impl From for ConfigFieldMapping { + fn from(value: feature_schema::ConfigFieldMapping) -> Self { + Self { + field_name: value.field_name, + field_number: value.field_number, + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct PeerConnStats { + pub rx_bytes: i64, + pub tx_bytes: i64, + pub rx_packets: i64, + pub tx_packets: i64, + pub latency_us: i64, +} + +impl From for PeerConnStats { + fn from(value: kernel_types::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, + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct PeerConnInfo { + pub conn_id: String, + pub my_peer_id: i64, + pub peer_id: i64, + pub features: Vec, + pub tunnel_type: Option, + pub local_addr: Option, + pub remote_addr: Option, + pub resolved_remote_addr: Option, + pub stats: Option, + pub loss_rate: Option, + pub is_client: bool, + pub network_name: Option, + pub is_closed: bool, + pub secure_auth_level: Option, + pub peer_identity_type: Option, +} + +impl From for PeerConnInfo { + fn from(value: kernel_types::PeerConnInfo) -> Self { + Self { + conn_id: value.conn_id, + my_peer_id: value.my_peer_id, + peer_id: value.peer_id, + features: value.features, + tunnel_type: value.tunnel_type, + local_addr: value.local_addr, + remote_addr: value.remote_addr, + resolved_remote_addr: value.resolved_remote_addr, + 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, + secure_auth_level: value.secure_auth_level, + peer_identity_type: value.peer_identity_type, + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct PeerInfo { + pub peer_id: i64, + pub default_conn_id: Option, + pub directly_connected_conns: Vec, + pub conns: Vec, +} + +impl From for PeerInfo { + fn from(value: kernel_types::PeerInfo) -> Self { + Self { + peer_id: value.peer_id, + default_conn_id: value.default_conn_id, + directly_connected_conns: value.directly_connected_conns, + conns: value.conns.into_iter().map(Into::into).collect(), + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct RouteView { + pub peer_id: i64, + pub hostname: Option, + pub ipv4: Option, + pub ipv4_cidr: Option, + pub ipv6_cidr: Option, + pub proxy_cidrs: Vec, + pub next_hop_peer_id: Option, + pub cost: Option, + pub path_latency: Option, + pub udp_nat_type: Option, + pub tcp_nat_type: Option, + pub inst_id: Option, + pub version: Option, + pub is_public_server: Option, +} + +impl From for RouteView { + fn from(value: kernel_types::RouteView) -> Self { + Self { + peer_id: value.peer_id, + hostname: value.hostname, + ipv4: value.ipv4, + ipv4_cidr: value.ipv4_cidr, + ipv6_cidr: value.ipv6_cidr, + proxy_cidrs: value.proxy_cidrs, + next_hop_peer_id: value.next_hop_peer_id, + cost: value.cost, + path_latency: value.path_latency, + udp_nat_type: value.udp_nat_type, + tcp_nat_type: value.tcp_nat_type, + inst_id: value.inst_id, + version: value.version, + is_public_server: value.is_public_server, + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct MyNodeInfo { + pub virtual_ipv4: Option, + pub virtual_ipv4_cidr: Option, + pub hostname: Option, + pub version: Option, + pub peer_id: Option, + pub listeners: Vec, + pub vpn_portal_cfg: Option, + pub udp_nat_type: Option, + pub tcp_nat_type: Option, +} + +impl From for MyNodeInfo { + fn from(value: kernel_types::MyNodeInfo) -> Self { + Self { + virtual_ipv4: value.virtual_ipv4, + virtual_ipv4_cidr: value.virtual_ipv4_cidr, + hostname: value.hostname, + version: value.version, + peer_id: value.peer_id, + listeners: value.listeners, + vpn_portal_cfg: value.vpn_portal_cfg, + udp_nat_type: value.udp_nat_type, + tcp_nat_type: value.tcp_nat_type, + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct RuntimeInstanceState { + pub config_id: String, + pub instance_id: String, + pub display_name: String, + pub running: bool, + pub tun_required: bool, + pub tun_attached: bool, + pub magic_dns_enabled: bool, + pub need_exit_node: bool, + pub error_message: Option, + pub my_node_info: Option, + pub events: Vec, + pub routes: Vec, + pub peers: Vec, +} + +impl From for RuntimeInstanceState { + fn from(value: kernel_types::RuntimeInstanceState) -> Self { + Self { + config_id: value.config_id, + instance_id: value.instance_id, + display_name: value.display_name, + running: value.running, + tun_required: value.tun_required, + tun_attached: value.tun_attached, + magic_dns_enabled: value.magic_dns_enabled, + need_exit_node: value.need_exit_node, + error_message: value.error_message, + my_node_info: value.my_node_info.map(Into::into), + events: value.events, + routes: value.routes.into_iter().map(Into::into).collect(), + peers: value.peers.into_iter().map(Into::into).collect(), + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct TunAggregateState { + pub active: bool, + pub attached_instance_ids: Vec, + pub aggregated_routes: Vec, + pub dns_servers: Vec, + pub need_rebuild: bool, +} + +impl From for TunAggregateState { + fn from(value: kernel_types::TunAggregateState) -> Self { + Self { + active: value.active, + attached_instance_ids: value.attached_instance_ids, + aggregated_routes: value.aggregated_routes, + dns_servers: value.dns_servers, + need_rebuild: value.need_rebuild, + } + } +} + +#[derive(Serialize)] +#[serde(rename_all = "camelCase")] +#[napi(object)] +pub struct RuntimeAggregateState { + pub instances: Vec, + pub tun: TunAggregateState, + pub running_instance_count: i32, +} + +impl From for RuntimeAggregateState { + fn from(value: kernel_types::RuntimeAggregateState) -> Self { + Self { + instances: value.instances.into_iter().map(Into::into).collect(), + tun: value.tun.into(), + running_instance_count: value.running_instance_count, + } + } +} diff --git a/easytier-contrib/easytier-ohrs/src/runtime.rs b/easytier-contrib/easytier-ohrs/src/runtime.rs deleted file mode 100644 index 33e14d22..00000000 --- a/easytier-contrib/easytier-ohrs/src/runtime.rs +++ /dev/null @@ -1 +0,0 @@ -pub(crate) mod state; diff --git a/easytier-contrib/easytier-ohrs/src/runtime/state/mod.rs b/easytier-contrib/easytier-ohrs/src/runtime/state/mod.rs deleted file mode 100644 index f84ecc4a..00000000 --- a/easytier-contrib/easytier-ohrs/src/runtime/state/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub(crate) mod runtime_state; diff --git a/easytier-core/src/gateway/proxy/tcp_proxy_service.rs b/easytier-core/src/gateway/proxy/tcp_proxy_service.rs index 964c7b3e..223211f8 100644 --- a/easytier-core/src/gateway/proxy/tcp_proxy_service.rs +++ b/easytier-core/src/gateway/proxy/tcp_proxy_service.rs @@ -262,6 +262,7 @@ impl UdpBindOptions { UdpBindOptions::socks5() + .with_need_protect(false) .with_context(self.socket_context.clone().with_ip_version(IpVersion::V6)) .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( Ipv6Addr::UNSPECIFIED, @@ -289,6 +290,10 @@ mod tests { ); assert_eq!(options[0].context.socket_mark, context.socket_mark); assert_eq!(options[0].context.ip_version, IpVersion::V6); + assert!( + !options[0].need_protect, + "SOCKS5 association socket is inbound" + ); } #[tokio::test] diff --git a/easytier-core/src/host/dns.rs b/easytier-core/src/host/dns.rs index cca66c09..3f2f994a 100644 --- a/easytier-core/src/host/dns.rs +++ b/easytier-core/src/host/dns.rs @@ -47,6 +47,10 @@ pub trait DnsRecordResolver: Send + Sync + 'static { /// 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. +/// When the host uses VPN bypass, every underlying DNS socket (including TCP +/// fallback) must be protected before sending queries. A system resolver is only +/// suitable if the host can guarantee equivalent routing; errors must not fall +/// back to unprotected DNS traffic. pub trait HostDnsIo: Send + Sync + 'static { fn submit_resolve(&self, operation: HostOperationId, query: &DnsQuery) -> io::Result<()>; diff --git a/easytier-core/src/host/environment.rs b/easytier-core/src/host/environment.rs index 6f63cd68..b8e5a068 100644 --- a/easytier-core/src/host/environment.rs +++ b/easytier-core/src/host/environment.rs @@ -13,6 +13,9 @@ use super::socket::{HostOperationId, HostSocketRuntime}; /// 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. +/// Source-address probes describe the underlay route: if VPN bypass is active, +/// protect the probe socket before connect/source-address selection. Protection +/// acknowledgement belongs inside this host operation, not a later guest call. pub trait HostConnectorEnvironmentIo: Send + Sync + 'static { fn submit_local_addr_for_remote( &self, diff --git a/easytier-core/src/host/socket/factory.rs b/easytier-core/src/host/socket/factory.rs index b1924731..e78ef8a2 100644 --- a/easytier-core/src/host/socket/factory.rs +++ b/easytier-core/src/host/socket/factory.rs @@ -31,6 +31,9 @@ pub struct HostUdpBindResult { /// 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. +/// When bind options request socket protection, the host must complete and +/// acknowledge protection before connect, datagram I/O, or successful creation +/// completion. A protection error fails the creation operation. pub trait HostSocketFactoryIo: HostSocketIo { fn submit_tcp_connect( &self, @@ -388,7 +391,8 @@ mod tests { .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), + .with_reuse_port(true) + .with_need_protect(true), purpose: TcpSocketPurpose::ManualConnect, }; let task = tokio::spawn({ @@ -407,6 +411,7 @@ mod tests { panic!("operation is not TCP connect"); }; assert_eq!(submitted, &options); + assert!(submitted.bind.need_protect); } io.complete_tcp( @@ -436,6 +441,7 @@ mod tests { let (runtime, factory) = test_factory(io.clone()); let options = UdpBindOptions { context: crate::socket::SocketContext::default().with_socket_mark(Some(9)), + need_protect: false, local_addr: Some("[::]:11013".parse().unwrap()), bind_device: Some("host-device".to_owned()), reuse_addr: true, @@ -459,6 +465,7 @@ mod tests { panic!("operation is not UDP bind"); }; assert_eq!(submitted, &options); + assert!(!submitted.need_protect); } io.complete_udp( diff --git a/easytier-core/src/host/socket/listener.rs b/easytier-core/src/host/socket/listener.rs index 62053ad8..049a53c0 100644 --- a/easytier-core/src/host/socket/listener.rs +++ b/easytier-core/src/host/socket/listener.rs @@ -19,6 +19,9 @@ pub struct HostTcpBindResult { /// 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. +/// When bind options request socket protection, the host must complete and +/// acknowledge protection before bind/listen and before exposing each accepted +/// child. A protection error fails the corresponding bind or accept operation. pub trait HostTcpListenerIo: HostSocketIo { fn submit_tcp_bind( &self, diff --git a/easytier-core/src/socket/mod.rs b/easytier-core/src/socket/mod.rs index e651f00d..d39164df 100644 --- a/easytier-core/src/socket/mod.rs +++ b/easytier-core/src/socket/mod.rs @@ -10,6 +10,12 @@ pub mod ring; pub mod tcp; pub mod udp; +/// Keep Rust constructors and deserialization consistent: host sockets normally +/// need VPN bypass; local/TUN-facing endpoint constructors explicitly opt out. +pub(crate) const fn default_need_protect() -> bool { + true +} + use std::{fmt::Debug, sync::Arc}; use async_trait::async_trait; diff --git a/easytier-core/src/socket/tcp.rs b/easytier-core/src/socket/tcp.rs index 526fb3df..1256ec3e 100644 --- a/easytier-core/src/socket/tcp.rs +++ b/easytier-core/src/socket/tcp.rs @@ -55,6 +55,12 @@ pub enum TcpSocketPurpose { pub struct TcpBindOptions { #[serde(default)] pub context: SocketContext, + /// Request host VPN bypass during creation, before bind/connect/listen. + /// Hosts with a protection service must await its acknowledgement; local + /// listeners opt out explicitly. Defaults to true, including when omitted + /// from serialized options. Accepted children inherit this requirement. + #[serde(default = "super::default_need_protect")] + pub need_protect: bool, pub local_addr: Option, pub bind_device: Option, /// `None` delegates the platform default to the host socket adapter. @@ -67,6 +73,7 @@ impl TcpBindOptions { pub fn new() -> Self { Self { context: SocketContext::default(), + need_protect: super::default_need_protect(), local_addr: None, bind_device: None, reuse_addr: None, @@ -90,6 +97,11 @@ impl TcpBindOptions { self } + pub fn with_need_protect(mut self, need_protect: bool) -> Self { + self.need_protect = need_protect; + self + } + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { self.context.ip_version = ip_version; self @@ -250,28 +262,36 @@ impl TcpListenOptions { pub fn proxy_nat(local_addr: SocketAddr) -> Self { Self { - bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + bind: TcpBindOptions::default() + .with_need_protect(false) + .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)), + bind: TcpBindOptions::default() + .with_need_protect(false) + .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)), + bind: TcpBindOptions::default() + .with_need_protect(false) + .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)), + bind: TcpBindOptions::default() + .with_need_protect(false) + .with_local_addr(Some(local_addr)), purpose: TcpListenPurpose::PortLease, } } @@ -612,7 +632,9 @@ mod tests { assert_eq!( TcpListenOptions::proxy_nat(local_addr), TcpListenOptions { - bind: TcpBindOptions::default().with_local_addr(Some(local_addr)), + bind: TcpBindOptions::default() + .with_need_protect(false) + .with_local_addr(Some(local_addr)), purpose: TcpListenPurpose::ProxyNat, } ); @@ -633,6 +655,7 @@ mod tests { options, TcpBindOptions { context: SocketContext::default().with_socket_mark(Some(7)), + need_protect: true, local_addr: Some(local_addr), bind_device: Some("eth0".to_owned()), reuse_addr: Some(true), @@ -645,6 +668,37 @@ mod tests { #[test] fn tcp_bind_default_delegates_reuse_addr_policy_to_host() { assert_eq!(TcpBindOptions::default().reuse_addr, None); + assert!(TcpBindOptions::default().need_protect); + } + + #[test] + fn tcp_constructor_protection_defaults_match_endpoint_role() { + let remote = SocketAddr::from(([192, 0, 2, 1], 11010)); + let local = SocketAddr::from(([0, 0, 0, 0], 11010)); + + assert!(TcpConnectOptions::direct_connect(remote).bind.need_protect); + assert!(TcpConnectOptions::proxy_nat(remote).bind.need_protect); + assert!(TcpListenOptions::direct_connect(local).bind.need_protect); + assert!(TcpListenOptions::hole_punch(local).bind.need_protect); + assert!(TcpListenOptions::manual_connect(local).bind.need_protect); + assert!(!TcpListenOptions::proxy_nat(local).bind.need_protect); + assert!(!TcpListenOptions::socks5(local).bind.need_protect); + assert!(!TcpListenOptions::port_forward(local).bind.need_protect); + assert!(!TcpListenOptions::port_lease(local).bind.need_protect); + } + + #[test] + fn tcp_bind_serde_defaults_to_protected_and_preserves_opt_out() { + let options: TcpBindOptions = serde_json::from_str( + r#"{"local_addr":null,"bind_device":null,"reuse_addr":null,"reuse_port":false,"only_v6":false}"#, + ) + .unwrap(); + assert!(options.need_protect); + let local = TcpListenOptions::port_forward("0.0.0.0:15555".parse().unwrap()); + let restored: TcpListenOptions = + serde_json::from_str(&serde_json::to_string(&local).unwrap()).unwrap(); + assert_eq!(restored, local); + assert!(!restored.bind.need_protect); } #[tokio::test] diff --git a/easytier-core/src/socket/udp/tests.rs b/easytier-core/src/socket/udp/tests.rs index 9be9cc39..557854de 100644 --- a/easytier-core/src/socket/udp/tests.rs +++ b/easytier-core/src/socket/udp/tests.rs @@ -39,6 +39,7 @@ fn bind_options_constructors_describe_socket_purpose() { reuse_port: false, only_v6: false, purpose: UdpSocketPurpose::HolePunchControl, + need_protect: true, } ); assert_eq!( @@ -51,6 +52,7 @@ fn bind_options_constructors_describe_socket_purpose() { reuse_port: false, only_v6: false, purpose: UdpSocketPurpose::HolePunchCandidate, + need_protect: true, } ); assert_eq!( @@ -63,6 +65,7 @@ fn bind_options_constructors_describe_socket_purpose() { reuse_port: false, only_v6: false, purpose: UdpSocketPurpose::DirectConnect, + need_protect: true, } ); assert_eq!( @@ -75,6 +78,7 @@ fn bind_options_constructors_describe_socket_purpose() { reuse_port: false, only_v6: false, purpose: UdpSocketPurpose::PortBoundListener, + need_protect: true, } ); assert_eq!( @@ -87,6 +91,7 @@ fn bind_options_constructors_describe_socket_purpose() { reuse_port: false, only_v6: false, purpose: UdpSocketPurpose::Socks5, + need_protect: true, } ); assert_eq!( @@ -99,7 +104,7 @@ fn bind_options_constructors_describe_socket_purpose() { ); assert_eq!( UdpBindOptions::default(), - UdpBindOptions::hole_punch_control() + UdpBindOptions::hole_punch_control().with_need_protect(true) ); } @@ -123,6 +128,7 @@ fn session_connect_request_keeps_peer_scoped_udp_shape() { reuse_port: false, only_v6: false, purpose: UdpSocketPurpose::PortBoundListener, + need_protect: true, } ); } @@ -2045,6 +2051,7 @@ async fn v4_hole_punch_control_sender_uses_factory_socket() { factory.bind_options(), vec![ UdpBindOptions::hole_punch_control() + .with_need_protect(false) .with_context(context.with_ip_version(IpVersion::V4)) .with_local_addr(Some(SocketAddr::V4(SocketAddrV4::new( Ipv4Addr::LOCALHOST, @@ -2089,6 +2096,7 @@ async fn v6_hole_punch_control_sender_uses_factory_socket() { factory.bind_options(), vec![ UdpBindOptions::hole_punch_control() + .with_need_protect(false) .with_context(context.with_ip_version(IpVersion::V6)) .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( Ipv6Addr::LOCALHOST, diff --git a/easytier-core/src/socket/udp/virtual_socket.rs b/easytier-core/src/socket/udp/virtual_socket.rs index 4620f927..ccea4cac 100644 --- a/easytier-core/src/socket/udp/virtual_socket.rs +++ b/easytier-core/src/socket/udp/virtual_socket.rs @@ -141,6 +141,7 @@ where let socket = factory .bind_udp( UdpBindOptions::hole_punch_control() + .with_need_protect(false) .with_context(context.with_ip_version(IpVersion::V4)) .with_local_addr(Some(SocketAddr::V4(SocketAddrV4::new( Ipv4Addr::LOCALHOST, @@ -167,6 +168,7 @@ where let socket = factory .bind_udp( UdpBindOptions::hole_punch_control() + .with_need_protect(false) .with_context(context.with_ip_version(IpVersion::V6)) .with_local_addr(Some(SocketAddr::V6(SocketAddrV6::new( Ipv6Addr::LOCALHOST, @@ -199,6 +201,11 @@ pub enum UdpSocketPurpose { pub struct UdpBindOptions { #[serde(default)] pub context: SocketContext, + /// Request host VPN bypass before bind or datagram I/O. Protection must be + /// acknowledged inside creation, not emitted as a fire-and-forget event. + /// Defaults to true; local/TUN-facing endpoints explicitly opt out. + #[serde(default = "crate::socket::default_need_protect")] + pub need_protect: bool, pub local_addr: Option, pub bind_device: Option, pub reuse_addr: bool, @@ -211,6 +218,7 @@ impl UdpBindOptions { fn for_purpose(purpose: UdpSocketPurpose) -> Self { Self { context: SocketContext::default(), + need_protect: crate::socket::default_need_protect(), local_addr: None, bind_device: None, reuse_addr: false, @@ -252,11 +260,15 @@ impl UdpBindOptions { } pub fn port_forward(local_addr: SocketAddr) -> Self { - Self::for_purpose(UdpSocketPurpose::PortForward).with_local_addr(Some(local_addr)) + Self::for_purpose(UdpSocketPurpose::PortForward) + .with_need_protect(false) + .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)) + Self::for_purpose(UdpSocketPurpose::PortLease) + .with_need_protect(false) + .with_local_addr(Some(local_addr)) } pub fn with_local_addr(mut self, local_addr: Option) -> Self { @@ -274,6 +286,11 @@ impl UdpBindOptions { self } + pub fn with_need_protect(mut self, need_protect: bool) -> Self { + self.need_protect = need_protect; + self + } + pub fn with_ip_version(mut self, ip_version: IpVersion) -> Self { self.context.ip_version = ip_version; self @@ -302,7 +319,44 @@ impl UdpBindOptions { impl Default for UdpBindOptions { fn default() -> Self { - Self::hole_punch_control() + // Preserve the existing default purpose/socket setup; local-only + // control packets opt out explicitly at their loopback call sites. + Self::for_purpose(UdpSocketPurpose::HolePunchControl) + } +} + +#[cfg(test)] +mod option_tests { + use super::*; + + #[test] + fn udp_constructor_protection_defaults_match_endpoint_role() { + let local = SocketAddr::from(([0, 0, 0, 0], 11010)); + + assert!(UdpBindOptions::default().need_protect); + assert!(UdpBindOptions::hole_punch_control().need_protect); + assert!(UdpBindOptions::hole_punch_candidate().need_protect); + assert!(UdpBindOptions::direct_connect().need_protect); + assert!(UdpBindOptions::port_bound_listener(local).need_protect); + assert!(UdpBindOptions::proxy_nat().need_protect); + assert!(UdpBindOptions::stun_probe().need_protect); + assert!(UdpBindOptions::socks5().need_protect); + assert!(!UdpBindOptions::port_forward(local).need_protect); + assert!(!UdpBindOptions::port_lease(local).need_protect); + } + + #[test] + fn udp_bind_serde_defaults_to_protected_and_preserves_opt_out() { + let options: UdpBindOptions = serde_json::from_str( + r#"{"local_addr":null,"bind_device":null,"reuse_addr":false,"reuse_port":false,"only_v6":false,"purpose":"DirectConnect"}"#, + ) + .unwrap(); + assert!(options.need_protect); + let local = UdpBindOptions::port_forward("0.0.0.0:15555".parse().unwrap()); + let restored: UdpBindOptions = + serde_json::from_str(&serde_json::to_string(&local).unwrap()).unwrap(); + assert_eq!(restored, local); + assert!(!restored.need_protect); } } diff --git a/easytier-core/src/wasi/imports.rs b/easytier-core/src/wasi/imports.rs index d19e7275..c98f6b3c 100644 --- a/easytier-core/src/wasi/imports.rs +++ b/easytier-core/src/wasi/imports.rs @@ -114,30 +114,36 @@ unsafe extern "C" { pub(crate) fn take_udp_send_ready(operation: u64) -> i32; /// Starts a TCP connection using an encoded `TcpConnectOptions` document. + /// Requested socket protection must complete before the host connects. 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. + /// Starts a UDP bind using an encoded `UdpBindOptions` document. Requested + /// socket protection must complete before bind or datagram I/O. 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. + /// Requested socket protection must complete before bind/listen. 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. + /// Starts accepting one TCP stream from a listener handle. A protected + /// listener's accepted child must be protected before it is exposed. 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`]. + /// When VPN bypass is active, the host must protect every underlying DNS + /// socket before sending a query or opening a DNS TCP connection. pub(crate) fn start_dns_resolve(operation: u64, query: u32, query_len: u32) -> i32; /// Probes or copies the encoded DNS address result for `operation`. @@ -159,6 +165,8 @@ unsafe extern "C" { 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`. + /// When VPN bypass is active, the host must protect the underlying route- + /// probe socket before connecting or sending through it. pub(crate) fn start_local_addr_for_remote( operation: u64, remote_addr: u32, diff --git a/easytier-core/src/wasi/wire/options.rs b/easytier-core/src/wasi/wire/options.rs index e933427f..1b5254a5 100644 --- a/easytier-core/src/wasi/wire/options.rs +++ b/easytier-core/src/wasi/wire/options.rs @@ -13,13 +13,13 @@ use crate::socket::{ use super::socket::{SOCKET_ADDRESS_LEN, decode_socket_address, encode_socket_address}; -const OPTIONS_VERSION: u8 = 2; +const OPTIONS_VERSION: u8 = 3; 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) + 76 + context_variable_len(&options.bind.context) + bind_device_len(&options.bind.bind_device), ); encoded.push(OPTIONS_VERSION); @@ -44,13 +44,14 @@ pub(crate) fn encode_tcp_connect_options(options: &TcpConnectOptions) -> io::Res TcpSocketPurpose::PortForward => 7, TcpSocketPurpose::DataPlane => 8, }); + encoded.push(u8::from(options.bind.need_protect)); 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), + 49 + context_variable_len(&options.context) + bind_device_len(&options.bind_device), ); encoded.push(OPTIONS_VERSION); encode_optional_address(&mut encoded, options.local_addr); @@ -69,13 +70,14 @@ pub(crate) fn encode_udp_bind_options(options: &UdpBindOptions) -> io::Result 7, UdpSocketPurpose::PortLease => 8, }); + encoded.push(u8::from(options.need_protect)); 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) + 49 + context_variable_len(&options.bind.context) + bind_device_len(&options.bind.bind_device), ); encoded.push(OPTIONS_VERSION); @@ -97,6 +99,7 @@ pub(crate) fn encode_tcp_listen_options(options: &TcpListenOptions) -> io::Resul TcpListenPurpose::PortForward => 5, TcpListenPurpose::PortLease => 6, }); + encoded.push(u8::from(options.bind.need_protect)); encode_bind_device(&mut encoded, &options.bind.bind_device)?; Ok(encoded) } @@ -221,13 +224,13 @@ mod tests { purpose: TcpSocketPurpose::ManualConnect, }; let encoded = encode_tcp_connect_options(&options).unwrap(); - assert_eq!(encoded.len(), 82); + assert_eq!(encoded.len(), 83); 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"); + assert_eq!(&encoded[66..71], &[2, 1, 1, 3, 1]); + assert_eq!(encoded[71], 1); + assert_eq!(&encoded[72..76], &7_u32.to_be_bytes()); + assert_eq!(&encoded[76..], b"device0"); } #[test] @@ -239,16 +242,16 @@ mod tests { .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]); + assert_eq!(&none[71..76], &[0, 0, 0, 0, 0]); + assert_eq!(&empty[71..76], &[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]); + assert_eq!(&udp_none[44..49], &[0, 0, 0, 0, 0]); + assert_eq!(&udp_empty[44..49], &[1, 0, 0, 0, 0]); let listen_none = encode_tcp_listen_options(&TcpListenOptions::direct_connect( "192.0.2.1:11013".parse().unwrap(), @@ -259,8 +262,8 @@ mod tests { .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]); + assert_eq!(&listen_none[44..49], &[0, 0, 0, 0, 0]); + assert_eq!(&listen_empty[44..49], &[1, 0, 0, 0, 0]); } #[test] @@ -329,6 +332,34 @@ mod tests { assert_eq!(udp[42], 5); } + #[test] + fn encodes_socket_protection_request_in_existing_options() { + let remote = "192.0.2.2:11013".parse().unwrap(); + let local = "0.0.0.0:11013".parse().unwrap(); + + let tcp_true = + encode_tcp_connect_options(&TcpConnectOptions::direct_connect(remote)).unwrap(); + let tcp_false = encode_tcp_connect_options( + &TcpConnectOptions::direct_connect(remote) + .with_bind(TcpBindOptions::default().with_need_protect(false)), + ) + .unwrap(); + assert_eq!(tcp_true[70], 1); + assert_eq!(tcp_false[70], 0); + + let listen_true = + encode_tcp_listen_options(&TcpListenOptions::direct_connect(local)).unwrap(); + let listen_false = encode_tcp_listen_options(&TcpListenOptions::proxy_nat(local)).unwrap(); + assert_eq!(listen_true[43], 1); + assert_eq!(listen_false[43], 0); + + let udp_true = encode_udp_bind_options(&UdpBindOptions::direct_connect()).unwrap(); + let udp_false = + encode_udp_bind_options(&UdpBindOptions::socks5().with_need_protect(false)).unwrap(); + assert_eq!(udp_true[43], 1); + assert_eq!(udp_false[43], 0); + } + #[test] fn encodes_gateway_purposes_with_stable_values() { let remote = "192.0.2.2:443".parse().unwrap(); diff --git a/easytier/src/common/dns.rs b/easytier/src/common/dns.rs index ede082a7..1c5fbe73 100644 --- a/easytier/src/common/dns.rs +++ b/easytier/src/common/dns.rs @@ -22,14 +22,19 @@ use hickory_resolver::{Resolver, TokioResolver}; use once_cell::sync::Lazy; use tokio::net::lookup_host; #[cfg(feature = "dns-resolver")] -use tokio::net::{TcpSocket, TcpStream, UdpSocket}; +use tokio::net::{TcpStream, UdpSocket}; #[cfg(feature = "dns-resolver")] use tokio::sync::Semaphore; use super::error::Error; use super::netns::NetNS; #[cfg(feature = "dns-resolver")] -use crate::tunnel::common::apply_socket_mark; +use crate::{ + socket::{tcp::create_tcp_socket, udp::create_udp_socket}, + socket_protector::native_socket_protection_available, +}; +#[cfg(feature = "dns-resolver")] +use easytier_core::socket::{NetNamespace, tcp::TcpBindOptions, udp::UdpBindOptions}; #[cfg(feature = "dns-resolver")] pub fn get_default_resolver_config() -> ResolverConfig { @@ -164,7 +169,14 @@ impl RuntimeDnsIoContext { #[cfg(feature = "dns-resolver")] fn is_process_default(&self) -> bool { - self.netns.is_none() && self.socket_mark.is_none() + self.netns.is_none() && self.socket_mark.is_none() && !native_socket_protection_available() + } + + #[cfg(feature = "dns-resolver")] + fn socket_context(&self) -> SocketContext { + SocketContext::default() + .with_netns(self.netns.clone().map(NetNamespace::new)) + .with_socket_mark(self.socket_mark) } } @@ -185,47 +197,6 @@ impl RuntimeDnsIoProvider { } } -#[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; @@ -243,13 +214,21 @@ impl RuntimeProvider for RuntimeDnsIoProvider { 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); + let options = TcpBindOptions::default() + .with_context(self.context.socket_context()) + .with_local_addr(bind_addr) + .with_bind_device(Some(String::new())) + .with_reuse_addr(false); 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 { + let connect = async { + let socket = create_tcp_socket(server_addr, &options) + .await + .map_err(io::Error::other)?; + socket.set_nodelay(true)?; + socket.connect(server_addr).await + }; + match tokio::time::timeout(wait_for, connect).await { Ok(Ok(stream)) => Ok(AsyncIoTokioAsStd(stream)), Ok(Err(error)) => Err(error), Err(_) => Err(io::Error::new( @@ -265,10 +244,10 @@ impl RuntimeProvider for RuntimeDnsIoProvider { 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 }) + let options = UdpBindOptions::default() + .with_context(self.context.socket_context()) + .with_local_addr(Some(local_addr)); + Box::pin(async move { create_udp_socket(&options).await.map_err(io::Error::other) }) } } @@ -339,7 +318,7 @@ impl RuntimeDnsResolver { context: RuntimeDnsIoContext, host: String, ) -> anyhow::Result> { - if context.socket_mark.is_some() { + if context.socket_mark.is_some() || native_socket_protection_available() { return Self::resolve_contextual_with_hickory(context, host).await; } @@ -381,8 +360,10 @@ impl DnsResolver for RuntimeDnsResolver { } #[cfg(not(feature = "dns-resolver"))] { - if context.socket_mark.is_some() { - anyhow::bail!("socket-marked DNS requires DNS resolver support"); + if context.socket_mark.is_some() + || crate::socket_protector::native_socket_protection_available() + { + anyhow::bail!("socket-marked or VPN-protected DNS requires DNS resolver support"); } if context.netns.is_none() { return Ok(resolve_ips(&query.host).await?); diff --git a/easytier/src/common/upnp.rs b/easytier/src/common/upnp.rs index de245d53..d53b268e 100644 --- a/easytier/src/common/upnp.rs +++ b/easytier/src/common/upnp.rs @@ -11,8 +11,10 @@ use easytier_core::connectivity::hole_punch::port_mapping::{ UdpPortMappingBackend, UdpPortMappingLifecycle, }; use natpmp::{ - Protocol as NatPmpProtocol, Response as NatPmpResponse, new_tokio_natpmp, new_tokio_natpmp_with, + Protocol as NatPmpProtocol, Response as NatPmpResponse, get_default_gateway, + new_natpmp_async_with, }; +use tokio::net::UdpSocket; use crate::igd_next::{ AddAnyPortError, Gateway, PortMappingProtocol, SearchOptions, search_gateway, @@ -26,6 +28,26 @@ const NAT_PMP_RESPONSE_TIMEOUT: Duration = Duration::from_secs(1); const UPNP_LEASE_DURATION_SECS: u32 = 300; const UPNP_DESCRIPTION: &str = "EasyTier udp hole punch"; +async fn new_protected_natpmp( + gateway: Option, +) -> anyhow::Result> { + let gateway = gateway + .map(Ok) + .unwrap_or_else(|| get_default_gateway().map_err(anyhow::Error::from))?; + let gateway_addr = SocketAddr::V4(SocketAddrV4::new(gateway, natpmp::NATPMP_PORT)); + let socket = crate::socket::udp::create_udp_socket( + &easytier_core::socket::udp::UdpBindOptions::direct_connect() + .with_local_addr(Some("0.0.0.0:0".parse().unwrap())), + ) + .await + .context("create protected nat-pmp socket")?; + socket + .connect(gateway_addr) + .await + .with_context(|| format!("connect nat-pmp socket to gateway {gateway}"))?; + Ok(new_natpmp_async_with(socket, gateway)) +} + enum PortMappingBackend { NatPmp { gateway: Ipv4Addr }, Igd { gateway: Gateway }, @@ -60,7 +82,9 @@ impl ActiveUdpPortMapping { async fn discover_nat_pmp_gateway( local_listener: &url::Url, ) -> anyhow::Result<(Ipv4Addr, SocketAddr)> { - let client = new_tokio_natpmp().await.context("create nat-pmp client")?; + let client = new_protected_natpmp(None) + .await + .context("create nat-pmp client")?; let gateway = *client.gateway(); let gateway_addr = SocketAddr::V4(SocketAddrV4::new(gateway, natpmp::NATPMP_PORT)); let local_addr = resolve_internal_addr(gateway_addr, local_listener).await?; @@ -428,7 +452,7 @@ async fn request_nat_pmp_mapping( public_port: u16, lifetime_secs: u32, ) -> anyhow::Result { - let client = new_tokio_natpmp_with(gateway) + let client = new_protected_natpmp(Some(gateway)) .await .with_context(|| format!("create nat-pmp client for gateway {gateway}"))?; client @@ -511,9 +535,13 @@ async fn resolve_internal_addr( listener_ipv4_host(local_listener).ok_or_else(|| anyhow!("listener must be ipv4"))?; let ip = if host.is_unspecified() { - let udp = std::net::UdpSocket::bind("0.0.0.0:0") - .context("bind probe socket for gateway route")?; + let options = easytier_core::socket::udp::UdpBindOptions::default() + .with_local_addr(Some("0.0.0.0:0".parse().unwrap())); + let udp = crate::socket::udp::create_udp_socket(&options) + .await + .context("create protected probe socket for gateway route")?; udp.connect(gateway_addr) + .await .with_context(|| format!("connect probe socket to gateway {gateway_addr}"))?; let SocketAddr::V4(local_addr) = udp.local_addr().context("get probe socket local addr")? else { diff --git a/easytier/src/host_runtime.rs b/easytier/src/host_runtime.rs index 3491dc3e..f3cfad07 100644 --- a/easytier/src/host_runtime.rs +++ b/easytier/src/host_runtime.rs @@ -108,25 +108,18 @@ impl ConnectorRuntime for NativeHostRuntime { 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)?; + let bind_addr = match remote_addr { + SocketAddr::V4(_) => { + SocketAddr::V4(SocketAddrV4::new(std::net::Ipv4Addr::UNSPECIFIED, 0)) + } + SocketAddr::V6(_) => { + SocketAddr::V6(SocketAddrV6::new(std::net::Ipv6Addr::UNSPECIFIED, 0, 0, 0)) + } + }; + let options = UdpBindOptions::default() + .with_context(context) + .with_local_addr(Some(bind_addr)); + let socket = crate::socket::udp::create_udp_socket(&options).await?; socket.connect(remote_addr).await?; Ok(socket.local_addr()?) } @@ -168,7 +161,9 @@ 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)?)) + Ok(Arc::new( + crate::socket::tcp::bind_tcp_listener(options).await?, + )) } } diff --git a/easytier/src/lib.rs b/easytier/src/lib.rs index a8634ac4..f7823951 100644 --- a/easytier/src/lib.rs +++ b/easytier/src/lib.rs @@ -20,6 +20,7 @@ pub mod rpc_service; #[cfg(feature = "management")] pub mod service_manager; pub(crate) mod socket; +pub mod socket_protector; pub mod tunnel; pub mod utils; #[cfg(feature = "web-client")] diff --git a/easytier/src/proto/rpc/standalone.rs b/easytier/src/proto/rpc/standalone.rs index 372771e6..9341f044 100644 --- a/easytier/src/proto/rpc/standalone.rs +++ b/easytier/src/proto/rpc/standalone.rs @@ -4,6 +4,7 @@ use easytier_core::{ }, rpc::standalone::StandAloneClient, socket::SocketListener, + socket::tcp::TcpBindOptions, socket::udp::{UdpBindOptions, UdpSessionListenRequest}, tunnel::Tunnel, }; @@ -28,7 +29,12 @@ pub fn runtime_rpc_client(remote_url: url::Url) -> RuntimeRpcClient { } pub fn runtime_rpc_listener(local_addr: std::net::SocketAddr) -> RuntimeRpcListener { - TcpTunnelListener::new(local_addr, native_host_runtime()) + // RPC clients may arrive through TUN; their replies must retain VPN routing. + let bind = TcpBindOptions::default() + .with_local_addr(Some(local_addr)) + .with_only_v6(true) + .with_need_protect(false); + TcpTunnelListener::new_with_bind(local_addr, bind, native_host_runtime()) } pub fn runtime_udp_tunnel_dialer(remote_url: url::Url) -> impl TunnelDialer { diff --git a/easytier/src/socket/tcp.rs b/easytier/src/socket/tcp.rs index db9709d9..3fb385d0 100644 --- a/easytier/src/socket/tcp.rs +++ b/easytier/src/socket/tcp.rs @@ -23,6 +23,7 @@ use tokio::{ use crate::{ common::netns::NetNS, + socket_protector::protect_native_socket, tunnel::common::{BindDev, apply_socket_mark, bind}, }; @@ -183,12 +184,7 @@ impl VirtualTcpSocket for RuntimeTcpSocket { pub struct RuntimeTcpListener { listener: TcpListener, purpose: TcpListenPurpose, -} - -impl RuntimeTcpListener { - pub(crate) fn new(listener: TcpListener, purpose: TcpListenPurpose) -> Self { - Self { listener, purpose } - } + need_protect: bool, } #[async_trait::async_trait] @@ -201,6 +197,7 @@ impl VirtualTcpListener for RuntimeTcpListener { async fn accept(&self) -> io::Result<(Self::Socket, SocketAddr)> { let (stream, addr) = self.listener.accept().await?; + protect_native_socket(&SockRef::from(&stream), self.need_protect).await?; if self.purpose == TcpListenPurpose::ProxyNat { prepare_proxy_tcp_socket(&stream)?; } @@ -229,45 +226,53 @@ fn bind_dev_from_options(options: &TcpBindOptions, local_addr_was_defaulted: boo }) } -fn bind_tcp_socket( +async fn bind_tcp_socket( remote_addr: SocketAddr, - bind_options: TcpBindOptions, + 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); + 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_addr(native_reuse_addr(bind_options)) .reuse_port(bind_options.reuse_port) .maybe_socket_mark(bind_options.context.socket_mark) + .need_protect(bind_options.need_protect) .call() + .await } -fn create_tcp_socket( +pub(crate) async fn create_tcp_socket( remote_addr: SocketAddr, bind_options: &TcpBindOptions, ) -> Result { + if must_bind_before_connect(bind_options) { + return bind_tcp_socket(remote_addr, bind_options).await; + } // 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) - }) + let socket = + NetNS::from_socket_context(&bind_options.context).run(|| -> Result<_, TunnelError> { + 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) + })?; + protect_native_socket(&SockRef::from(&socket), bind_options.need_protect).await?; + Ok(socket) } fn must_bind_before_connect(bind_options: &TcpBindOptions) -> bool { @@ -290,30 +295,23 @@ fn native_reuse_addr(bind_options: &TcpBindOptions) -> bool { .unwrap_or_else(native_reuse_addr_default) } -pub(crate) fn bind_tcp_listener( +pub(crate) async 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 mut bind_options = options.bind; 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)) + if bind_options.bind_device.is_none() && purpose == TcpListenPurpose::PortLease { + bind_options.bind_device = Some(String::new()); + } + let listener = create_tcp_socket(addr, &bind_options).await?.listen(1024)?; + Ok(RuntimeTcpListener { + listener, + purpose, + need_protect: bind_options.need_protect, + }) } pub(crate) async fn connect_tcp( @@ -323,14 +321,7 @@ pub(crate) async fn connect_tcp( 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 socket = create_tcp_socket(remote_addr, &bind_options).await?; let stream = socket.connect(remote_addr).await?; prepare_connected_tcp_socket(&stream, purpose)?; Ok(RuntimeTcpSocket::new(stream)) @@ -398,6 +389,12 @@ mod tests { #[test] fn tcp_connect_binds_when_socket_option_requires_pre_connect_setup() { + assert!(must_bind_before_connect( + &TcpBindOptions::default().with_local_addr(Some("127.0.0.1:0".parse().unwrap())) + )); + assert!(must_bind_before_connect( + &TcpBindOptions::default().with_reuse_port(true) + )); assert!(must_bind_before_connect( &TcpBindOptions::default().with_only_v6(true) )); diff --git a/easytier/src/socket/udp.rs b/easytier/src/socket/udp.rs index 9e2e2cbb..0f413c78 100644 --- a/easytier/src/socket/udp.rs +++ b/easytier/src/socket/udp.rs @@ -167,28 +167,26 @@ impl RuntimeUdpSocketFactory { | 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, - ))) - } +/// Creates a host-owned UDP socket, including any required protection, before +/// it can be used by DNS, a route probe, or a virtual socket adapter. +pub(crate) async fn create_udp_socket(options: &UdpBindOptions) -> anyhow::Result { + let factory = RuntimeUdpSocketFactory::new(); + let bind_addr = options + .local_addr + .unwrap_or_else(|| SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0))); + Ok(bind::() + .addr(bind_addr) + .dev(factory.bind_device_for(options)) + .net_ns(NetNS::from_socket_context(&options.context)) + .only_v6(options.only_v6) + .reuse_addr(factory.reuse_addr_for(options)) + .reuse_port(options.reuse_port) + .maybe_socket_mark(options.context.socket_mark) + .need_protect(options.need_protect) + .call() + .await?) } #[async_trait] @@ -196,7 +194,11 @@ impl VirtualUdpSocketFactory for RuntimeUdpSocketFactory { type Socket = RuntimeUdpSocket; async fn bind_udp(&self, options: UdpBindOptions) -> anyhow::Result> { - self.bind_udp_socket(options) + let socket = create_udp_socket(&options).await?; + Ok(Arc::new(RuntimeUdpSocket::new_with_context( + Arc::new(socket), + options.context, + ))) } } diff --git a/easytier/src/socket_protector.rs b/easytier/src/socket_protector.rs new file mode 100644 index 00000000..08065333 --- /dev/null +++ b/easytier/src/socket_protector.rs @@ -0,0 +1,235 @@ +use std::{ + io, + sync::{Arc, RwLock}, +}; + +use async_trait::async_trait; + +#[cfg(unix)] +use std::os::fd::AsRawFd; +#[cfg(windows)] +use std::os::windows::io::AsRawSocket; + +/// Native-only platform callback used by socket creation when `need_protect` +/// is set. The future must resolve only after protection is actually applied; +/// emitting an event without awaiting its acknowledgement is not sufficient. +/// Errors fail creation before bind/connect/listen or exposing an accepted child. +/// +/// WASI embedders implement the same contract inside their existing host socket +/// creation operations, using core's encoded bind options, not this raw-FD API. +#[async_trait] +pub trait NativeSocketProtector: Send + Sync + 'static { + async fn protect(&self, socket_handle: u64) -> io::Result<()>; +} + +static NATIVE_SOCKET_PROTECTOR: RwLock>> = RwLock::new(None); + +/// Installs or removes the process-wide native socket protection capability. +/// +/// Instance-specific routing policy still travels in socket options; this hook +/// only exposes a platform service such as Android/iOS/HarmonyOS VPN bypass. +pub fn set_native_socket_protector(protector: Option>) { + let mut guard = NATIVE_SOCKET_PROTECTOR + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner); + *guard = protector; +} + +fn native_socket_protector() -> Option> { + #[cfg(all(test, unix))] + if let Ok(protector) = TEST_SOCKET_PROTECTOR.try_with(Arc::clone) { + return Some(protector); + } + NATIVE_SOCKET_PROTECTOR + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone() +} + +// Isolate factory ordering tests from other tests' sockets in the same process. +#[cfg(all(test, unix))] +tokio::task_local! { + static TEST_SOCKET_PROTECTOR: Arc; +} + +pub(crate) fn native_socket_protection_available() -> bool { + native_socket_protector().is_some() +} + +// socket2 owns the platform-specific handle; SockRef adapts Tokio sockets here. +pub(crate) async fn protect_native_socket( + socket: &socket2::Socket, + need_protect: bool, +) -> io::Result<()> { + if !need_protect { + return Ok(()); + } + let Some(protector) = native_socket_protector() else { + return Ok(()); + }; + #[cfg(unix)] + let handle = u64::try_from(socket.as_raw_fd()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "invalid socket fd"))?; + #[cfg(windows)] + let handle = socket.as_raw_socket() as u64; + protector.protect(handle).await +} + +#[cfg(all(test, unix))] +mod tests { + use super::*; + use easytier_core::socket::{ + SocketListener, + tcp::{TcpBindOptions, TcpConnectOptions, TcpListenOptions, VirtualTcpListener}, + udp::UdpBindOptions, + }; + use std::{ + os::fd::BorrowedFd, + sync::atomic::{AtomicUsize, Ordering}, + time::Duration, + }; + use tokio::sync::Semaphore; + + struct GateProtector { + calls: AtomicUsize, + gate: Semaphore, + fail: bool, + expect_unbound: bool, + } + + impl GateProtector { + fn new(fail: bool, expect_unbound: bool) -> Arc { + Arc::new(Self { + calls: AtomicUsize::new(0), + gate: Semaphore::new(0), + fail, + expect_unbound, + }) + } + } + + #[async_trait] + impl NativeSocketProtector for GateProtector { + async fn protect(&self, handle: u64) -> io::Result<()> { + // The caller keeps the socket alive while this borrowed callback runs. + let fd = unsafe { BorrowedFd::borrow_raw(i32::try_from(handle).unwrap()) }; + let socket = socket2::SockRef::from(&fd); + if self.expect_unbound || self.calls.load(Ordering::SeqCst) == 0 { + assert_eq!(socket.local_addr()?.as_socket().unwrap().port(), 0); + assert!( + socket.peer_addr().is_err(), + "connect must wait for protect completion" + ); + } + self.calls.fetch_add(1, Ordering::SeqCst); + if self.fail { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "test protection failure", + )); + } + self.gate.acquire().await.unwrap().forget(); + Ok(()) + } + } + + #[tokio::test] + async fn tcp_connect_waits_for_protection_ack() { + let protector = GateProtector::new(false, true); + TEST_SOCKET_PROTECTOR + .scope(protector.clone(), async { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let options = TcpConnectOptions::direct_connect(listener.local_addr().unwrap()); + let connect = crate::socket::tcp::connect_tcp(options); + tokio::pin!(connect); + assert!(futures::poll!(&mut connect).is_pending()); + assert_eq!(protector.calls.load(Ordering::SeqCst), 1); + assert!( + tokio::time::timeout(Duration::from_millis(20), listener.accept()) + .await + .is_err() + ); + protector.gate.add_permits(1); + let _client = connect.await.unwrap(); + listener.accept().await.unwrap(); + }) + .await; + } + + #[tokio::test] + async fn protection_failure_blocks_creation_and_explicit_false_bypasses_callback() { + let protector = GateProtector::new(true, true); + TEST_SOCKET_PROTECTOR + .scope(protector.clone(), async { + let local = "127.0.0.1:0".parse().unwrap(); + let bind = TcpBindOptions::default().with_local_addr(Some(local)); + assert!( + crate::socket::tcp::create_tcp_socket(local, &bind) + .await + .is_err() + ); + assert!( + crate::socket::tcp::bind_tcp_listener(TcpListenOptions::direct_connect(local)) + .await + .is_err() + ); + let udp = UdpBindOptions::direct_connect().with_local_addr(Some(local)); + assert!(crate::socket::udp::create_udp_socket(&udp).await.is_err()); + assert_eq!(protector.calls.load(Ordering::SeqCst), 3); + + for options in [ + TcpListenOptions::proxy_nat(local), + TcpListenOptions::socks5(local), + TcpListenOptions::port_forward(local), + TcpListenOptions::port_lease(local), + ] { + let listener = crate::socket::tcp::bind_tcp_listener(options) + .await + .unwrap(); + let connect = TcpConnectOptions::direct_connect(listener.local_addr().unwrap()) + .with_bind(TcpBindOptions::default().with_need_protect(false)); + let _client = crate::socket::tcp::connect_tcp(connect).await.unwrap(); + listener.accept().await.unwrap(); + } + let udp = udp.with_need_protect(false); + crate::socket::udp::create_udp_socket(&udp).await.unwrap(); + let mut rpc = crate::proto::rpc::standalone::runtime_rpc_listener(local); + rpc.listen().await.unwrap(); + let rpc_addr = rpc.local_url().socket_addrs(|| None).unwrap()[0]; + let _client = tokio::net::TcpStream::connect(rpc_addr).await.unwrap(); + rpc.accept().await.unwrap(); + assert_eq!(protector.calls.load(Ordering::SeqCst), 3); + }) + .await; + } + + #[tokio::test] + async fn accepted_child_waits_for_inherited_protection() { + let protector = GateProtector::new(false, false); + TEST_SOCKET_PROTECTOR + .scope(protector.clone(), async { + let bind = crate::socket::tcp::bind_tcp_listener(TcpListenOptions::direct_connect( + "127.0.0.1:0".parse().unwrap(), + )); + tokio::pin!(bind); + assert!(futures::poll!(&mut bind).is_pending()); + assert_eq!(protector.calls.load(Ordering::SeqCst), 1); + protector.gate.add_permits(1); + let listener = bind.await.unwrap(); + let _client = tokio::net::TcpStream::connect(listener.local_addr().unwrap()) + .await + .unwrap(); + let accept = listener.accept(); + tokio::pin!(accept); + assert!( + tokio::time::timeout(Duration::from_millis(20), &mut accept) + .await + .is_err() + ); + assert_eq!(protector.calls.load(Ordering::SeqCst), 2); + protector.gate.add_permits(1); + accept.await.unwrap(); + }) + .await; + } +} diff --git a/easytier/src/tunnel/common.rs b/easytier/src/tunnel/common.rs index bffebfda..4ead8827 100644 --- a/easytier/src/tunnel/common.rs +++ b/easytier/src/tunnel/common.rs @@ -72,7 +72,6 @@ fn setup_socket2_ext( only_v6: bool, reuse_addr: bool, reuse_port: bool, - socket_mark: Option, ) -> Result<(), TunnelError> { #[cfg(target_os = "windows")] { @@ -95,11 +94,6 @@ fn setup_socket2_ext( 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. - // Accepted child sockets inherit the mark from the listener on Linux. - apply_socket_mark(socket2_socket, socket_mark)?; - if let Err(e) = socket2_socket.bind(&socket2::SockAddr::from(*bind_addr)) { if bind_addr.is_ipv4() { return Err(e.into()); @@ -197,6 +191,8 @@ impl From<&str> for BindDev { /// This function creates a new socket, applies specific configurations (such as /// binding to a device or setting IPv6-only flags), and finalizes it into the /// requested [`Bindable`] type. +/// Protection is acknowledged before bind/listen; namespace guards never cross +/// an await. Failures drop the unbound socket instead of allowing early I/O. /// /// # Arguments /// @@ -216,7 +212,7 @@ impl From<&str> for BindDev { /// /// Returns a [`TunnelError`] if socket creation, configuration, or finalization fails. #[builder] -pub fn bind( +pub async fn bind( addr: SocketAddr, #[builder(default, into)] dev: BindDev, net_ns: Option, @@ -226,23 +222,26 @@ pub fn bind( /// 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, + #[builder(default)] need_protect: bool, ) -> Result { - let _g = net_ns.map(|n| n.guard()); - let dev = match dev { - BindDev::Auto => get_interface_name_by_ip(&addr.ip()), - BindDev::Disabled => None, - BindDev::Custom(s) => Some(s), + let (socket, dev) = { + let _g = net_ns.as_ref().map(|n| n.guard()); + let dev = match dev { + BindDev::Auto => get_interface_name_by_ip(&addr.ip()), + BindDev::Disabled => None, + BindDev::Custom(s) => Some(s), + }; + let socket = + socket2::Socket::new(socket2::Domain::for_address(addr), B::TYPE, B::PROTOCOL)?; + // A platform protector may set its own routing mark. Never overwrite + // that mark with caller options after its acknowledgement. + apply_socket_mark(&socket, socket_mark)?; + (socket, dev) }; - let socket = socket2::Socket::new(socket2::Domain::for_address(addr), B::TYPE, B::PROTOCOL)?; - setup_socket2_ext( - &socket, - &addr, - dev, - only_v6, - reuse_addr, - reuse_port, - socket_mark, - )?; + crate::socket_protector::protect_native_socket(&socket, need_protect).await?; + // Restore the requested namespace for synchronous bind/device lookup only. + let _g = net_ns.as_ref().map(|n| n.guard()); + setup_socket2_ext(&socket, &addr, dev, only_v6, reuse_addr, reuse_port)?; B::finalize(socket) } @@ -322,8 +321,8 @@ pub(crate) mod tests { target_os = "linux", target_env = "ohos" ))] - #[test] - fn bind_custom_device_is_applied_for_unspecified_addr() { + #[tokio::test] + async fn bind_custom_device_is_applied_for_unspecified_addr() { use std::net::SocketAddr; use tokio::net::UdpSocket; @@ -332,6 +331,7 @@ pub(crate) mod tests { .addr(addr) .dev("et/invalid-device-name") .call() + .await .expect_err("custom device must not be skipped for unspecified bind addr"); } diff --git a/easytier/src/tunnel/websocket.rs b/easytier/src/tunnel/websocket.rs index 96d4d9dd..df16da05 100644 --- a/easytier/src/tunnel/websocket.rs +++ b/easytier/src/tunnel/websocket.rs @@ -312,7 +312,8 @@ impl WsTunnelListener { .addr(addr) .only_v6(true) .maybe_socket_mark(self.socket_mark) - .call()?; + .call() + .await?; self.addr .set_port(Some(listener.local_addr()?.port()))