refactor(ohos): 拆分 OHRS 包并按 socket 精细保护 VPN 流量 (#2543)

* 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 <frankhan@FrankHans-Mac-mini.local>
Co-authored-by: KKRainbow <443152178@qq.com>
This commit is contained in:
authored and GitHub committed 2026-09-09 22:12:47 +08:00
1 parent b1f87f025b
commit 38e2a621bb
63 files changed
+2011 -397

No files matched your search

+7
View File
@@ -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')
+84
View File
@@ -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.
+35 -6
View File
@@ -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",
+9 -6
View File
@@ -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"
@@ -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.
@@ -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"
@@ -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<Runtime> = 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<Arc<NativeInstanceManager>> = 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(),
);
}
}
@@ -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<UnixStream>) {
}
}
pub(crate) fn broadcast_local_socket_message(
pub fn broadcast_local_socket_message(
clients: &mut Vec<UnixStream>,
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<UnixStream>,
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());
}
}
@@ -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<String>) -> Vec<String> {
.collect()
}
pub(crate) fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec<String> {
pub fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec<String> {
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<Strin
raw_routes.push(cidr);
}
raw_routes.extend(manual_routes.iter().cloned());
raw_routes.extend(config_proxy_cidrs.iter().cloned());
raw_routes.extend(instance.manual_routes.iter().cloned());
// Local proxy CIDRs are advertisements for networks reached through this
// node. Installing them into the same local TUN would recapture the proxy's
// own destination sockets instead of using the physical network.
raw_routes.extend(runtime_proxy_cidrs.iter().cloned());
simplify_routes(raw_routes)
}
pub(crate) fn aggregate_requested_tun_routes(instances: &[RuntimeInstanceState]) -> Vec<String> {
pub fn aggregate_requested_tun_routes(instances: &[RuntimeInstanceState]) -> Vec<String> {
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()));
}
}
@@ -0,0 +1 @@
pub mod state;
@@ -0,0 +1 @@
pub mod runtime_state;
@@ -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<String>,
@@ -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<String>,
@@ -91,7 +86,6 @@ pub struct RouteView {
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
#[napi(object)]
pub struct MyNodeInfo {
pub virtual_ipv4: Option<String>,
pub virtual_ipv4_cidr: Option<String>,
@@ -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<String>,
pub routes: Vec<RouteView>,
pub peers: Vec<PeerInfo>,
#[serde(skip)]
pub manual_routes: Vec<String>,
}
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
#[napi(object)]
pub struct TunAggregateState {
pub active: bool,
pub attached_instance_ids: Vec<String>,
@@ -136,7 +130,6 @@ pub struct TunAggregateState {
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
#[napi(object)]
pub struct RuntimeAggregateState {
pub instances: Vec<RuntimeInstanceState>,
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<api::manage::NetworkConfig>,
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,
}
}
@@ -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<SocketProtectionRequest>,
pending: HashMap<u64, PendingSocketProtection>,
}
struct PendingSocketProtection {
completion: oneshot::Sender<io::Result<()>>,
_socket: OwnedFd,
}
#[derive(Default)]
pub struct SocketProtectionManager {
state: Mutex<SocketProtectionState>,
request_ready: Notify,
}
pub static SOCKET_PROTECTION_MANAGER: Lazy<Arc<SocketProtectionManager>> =
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::<Vec<_>>();
queued_ids
.into_iter()
.filter_map(|request_id| state.pending.remove(&request_id))
.map(|pending| pending.completion)
.collect::<Vec<_>>()
};
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<SocketProtectionRequest> {
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<String>) -> 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());
});
}
}
@@ -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"] }
@@ -0,0 +1,4 @@
pub mod repository;
pub mod services;
pub mod storage;
pub mod types;
@@ -0,0 +1,2 @@
pub mod schema_service;
pub mod share_link_service;
@@ -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<FieldOption> {
.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");
}
}
@@ -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::<NetworkConfig>(&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);
}
}
@@ -0,0 +1 @@
pub mod config_meta;
@@ -0,0 +1 @@
pub mod stored_config;
@@ -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<StoredConfigMeta>,
}
#[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<String>,
@@ -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,
@@ -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<Option<PathBuf>> = Mutex::new(None);
static RUNTIME_CONFIG_SNAPSHOTS: Lazy<Mutex<HashMap<String, RuntimeConfigSnapshot>>> =
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<RuntimeConfigSnapshot> {
pub fn get_runtime_config_snapshot(config_id: &str) -> Option<RuntimeConfigSnapshot> {
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<String>, Vec<String>) {
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<PathBuf> {
pub fn config_root_dir() -> Option<PathBuf> {
CONFIG_ROOT_DIR
.lock()
.ok()
.and_then(|guard| guard.as_ref().cloned())
}
pub(crate) fn kernel_socket_path() -> Option<PathBuf> {
pub fn kernel_socket_path() -> Option<PathBuf> {
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<String> {
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<StoredConfigRecord> {
@@ -292,24 +281,6 @@ pub fn create_config_record(config_id: String, display_name: String) -> Option<S
save_config_record(config_id, display_name, normalized_json)
}
pub fn start_kernel_with_config_id(config_id: &str) -> 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::<NetworkConfig>(&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
);
}
@@ -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<FeatureLogSink> = 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(),
);
}
}
@@ -1,4 +0,0 @@
pub(crate) mod repository;
pub(crate) mod services;
pub(crate) mod storage;
pub(crate) mod types;
@@ -1,2 +0,0 @@
pub(crate) mod schema_service;
pub(crate) mod share_link_service;
@@ -1 +0,0 @@
pub(crate) mod config_meta;
@@ -1 +0,0 @@
pub(crate) mod stored_config;
@@ -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)
}
@@ -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(),
});
}
}
@@ -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};
@@ -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};
+76 -40
View File
@@ -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<Runtime> = 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<Arc<NativeInstanceManager>> =
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<Mutex<HashMap<String, ManagedWebClient>>> =
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<String> {
resolve_instance_id_from_state(&collect_runtime_state_inner(), instance_name)
}
pub(crate) fn build_default_network_config_json() -> Result<String, String> {
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<String, String> {
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::<NetworkConfig>(&raw) {
cache_runtime_config_snapshot(config_id.to_string(), display_name, config);
}
started
}
fn parse_instance_uuid(config_id: &str) -> Option<Uuid> {
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<KeyValuePair> {
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<SocketProtectionRequest> {
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<String>,
) -> bool {
let Ok(request_id) = request_id.parse::<u64>() 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<ConfigFieldMapping> {
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<bool>) -> O
#[napi]
pub fn parse_config_share_link(share_link: String) -> Option<SharedConfigLinkPayload> {
parse_config_share_link_inner(&share_link)
parse_config_share_link_inner(&share_link).map(Into::into)
}
#[napi]
@@ -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<easytier_ohos_core::socket_protection::SocketProtectionRequest>
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<feature_types::StoredConfigMeta> 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<feature_types::StoredConfigRecord> 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<StoredConfigMeta>,
}
impl From<feature_types::StoredConfigList> 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<feature_types::ExportTomlResult> 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<String>,
pub only_start: bool,
}
impl From<feature_types::SharedConfigLinkPayload> 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<feature_types::KeyValuePair> 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<feature_types::SnapshotImportResult> 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<feature_schema::FieldOption> 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<feature_schema::ValidationRule> 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<String>,
pub semantic_type: Option<String>,
pub value_kind: String,
pub is_list: bool,
pub required: bool,
pub default_value_text: Option<String>,
pub enum_options: Vec<FieldOption>,
pub validations: Vec<ValidationRule>,
pub children: Vec<NetworkConfigSchema>,
pub definitions: Vec<NetworkConfigSchema>,
}
impl From<feature_schema::NetworkConfigSchema> 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<feature_schema::ConfigFieldMapping> 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<kernel_types::PeerConnStats> 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<String>,
pub tunnel_type: Option<String>,
pub local_addr: Option<String>,
pub remote_addr: Option<String>,
pub resolved_remote_addr: Option<String>,
pub stats: Option<PeerConnStats>,
pub loss_rate: Option<f64>,
pub is_client: bool,
pub network_name: Option<String>,
pub is_closed: bool,
pub secure_auth_level: Option<i32>,
pub peer_identity_type: Option<i32>,
}
impl From<kernel_types::PeerConnInfo> 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<String>,
pub directly_connected_conns: Vec<String>,
pub conns: Vec<PeerConnInfo>,
}
impl From<kernel_types::PeerInfo> 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<String>,
pub ipv4: Option<String>,
pub ipv4_cidr: Option<String>,
pub ipv6_cidr: Option<String>,
pub proxy_cidrs: Vec<String>,
pub next_hop_peer_id: Option<i64>,
pub cost: Option<i32>,
pub path_latency: Option<i64>,
pub udp_nat_type: Option<i32>,
pub tcp_nat_type: Option<i32>,
pub inst_id: Option<String>,
pub version: Option<String>,
pub is_public_server: Option<bool>,
}
impl From<kernel_types::RouteView> 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<String>,
pub virtual_ipv4_cidr: Option<String>,
pub hostname: Option<String>,
pub version: Option<String>,
pub peer_id: Option<i64>,
pub listeners: Vec<String>,
pub vpn_portal_cfg: Option<String>,
pub udp_nat_type: Option<i32>,
pub tcp_nat_type: Option<i32>,
}
impl From<kernel_types::MyNodeInfo> 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<String>,
pub my_node_info: Option<MyNodeInfo>,
pub events: Vec<String>,
pub routes: Vec<RouteView>,
pub peers: Vec<PeerInfo>,
}
impl From<kernel_types::RuntimeInstanceState> 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<String>,
pub aggregated_routes: Vec<String>,
pub dns_servers: Vec<String>,
pub need_rebuild: bool,
}
impl From<kernel_types::TunAggregateState> 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<RuntimeInstanceState>,
pub tun: TunAggregateState,
pub running_instance_count: i32,
}
impl From<kernel_types::RuntimeAggregateState> 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,
}
}
}
@@ -1 +0,0 @@
pub(crate) mod state;
@@ -1 +0,0 @@
pub(crate) mod runtime_state;
@@ -262,6 +262,7 @@ impl<R: TcpProxyRuntime + 'static, F: VirtualTcpListenerFactory, C: TcpProxyDest
.bind_tcp(
TcpListenOptions::proxy_nat(listen_addr).with_bind(
TcpBindOptions::default()
.with_need_protect(false)
.with_context(
self.socket_context
.clone()
+5
View File
@@ -67,6 +67,7 @@ where
fn udp_bind_options(&self) -> 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]
+4
View File
@@ -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<()>;
+3
View File
@@ -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,
+8 -1
View File
@@ -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(
@@ -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,
+6
View File
@@ -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;
+59 -5
View File
@@ -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<SocketAddr>,
pub bind_device: Option<String>,
/// `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]
+9 -1
View File
@@ -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,
+57 -3
View File
@@ -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<SocketAddr>,
pub bind_device: Option<String>,
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<SocketAddr>) -> 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);
}
}
+10 -2
View File
@@ -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,
+46 -15
View File
@@ -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<Vec<u8>> {
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<Vec<u8>> {
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<Ve
UdpSocketPurpose::PortForward => 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<Vec<u8>> {
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();
+37 -56
View File
@@ -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<SocketAddr>,
) -> io::Result<TcpSocket> {
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<UdpSocket> {
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 = <TokioRuntimeProvider as RuntimeProvider>::Handle;
@@ -243,13 +214,21 @@ impl RuntimeProvider for RuntimeDnsIoProvider {
bind_addr: Option<SocketAddr>,
wait_for: Option<Duration>,
) -> Pin<Box<dyn Send + Future<Output = io::Result<Self::Tcp>>>> {
// 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<Box<dyn Send + Future<Output = io::Result<Self::Udp>>>> {
// 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<Vec<IpAddr>> {
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?);
+33 -5
View File
@@ -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<Ipv4Addr>,
) -> anyhow::Result<natpmp::NatpmpAsync<UdpSocket>> {
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<u16> {
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 {
+15 -20
View File
@@ -108,25 +108,18 @@ impl ConnectorRuntime for NativeHostRuntime {
remote_addr: SocketAddr,
context: SocketContext,
) -> anyhow::Result<SocketAddr> {
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<Arc<Self::Listener>> {
Ok(Arc::new(crate::socket::tcp::bind_tcp_listener(options)?))
Ok(Arc::new(
crate::socket::tcp::bind_tcp_listener(options).await?,
))
}
}
+1
View File
@@ -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")]
+7 -1
View File
@@ -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 {
+46 -49
View File
@@ -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<TcpSocket, TunnelError> {
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::<TcpSocket>()
.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<TcpSocket, TunnelError> {
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<RuntimeTcpListener, TunnelError> {
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::<TcpListener>()
.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)
));
+24 -22
View File
@@ -167,28 +167,26 @@ impl RuntimeUdpSocketFactory {
| UdpSocketPurpose::PortForward
) && !cfg!(target_os = "windows"))
}
}
fn bind_udp_socket(&self, options: UdpBindOptions) -> anyhow::Result<Arc<RuntimeUdpSocket>> {
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::<UdpSocket>()
.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<UdpSocket> {
let factory = RuntimeUdpSocketFactory::new();
let bind_addr = options
.local_addr
.unwrap_or_else(|| SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0)));
Ok(bind::<UdpSocket>()
.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<Arc<Self::Socket>> {
self.bind_udp_socket(options)
let socket = create_udp_socket(&options).await?;
Ok(Arc::new(RuntimeUdpSocket::new_with_context(
Arc::new(socket),
options.context,
)))
}
}
+235
View File
@@ -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<Option<Arc<dyn NativeSocketProtector>>> = 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<Arc<dyn NativeSocketProtector>>) {
let mut guard = NATIVE_SOCKET_PROTECTOR
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*guard = protector;
}
fn native_socket_protector() -> Option<Arc<dyn NativeSocketProtector>> {
#[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<dyn NativeSocketProtector>;
}
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<Self> {
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;
}
}
+24 -24
View File
@@ -72,7 +72,6 @@ fn setup_socket2_ext(
only_v6: bool,
reuse_addr: bool,
reuse_port: bool,
socket_mark: Option<u32>,
) -> 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<B: Bindable>(
pub async fn bind<B: Bindable>(
addr: SocketAddr,
#[builder(default, into)] dev: BindDev,
net_ns: Option<NetNS>,
@@ -226,23 +222,26 @@ pub fn bind<B: Bindable>(
/// 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<u32>,
#[builder(default)] need_protect: bool,
) -> Result<B, TunnelError> {
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");
}
+2 -1
View File
@@ -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()))