Compare commits

..
Author SHA1 Message Date
fanyang89 c12b73e1ae fix(stats): use age-based GC check to avoid evicting fresh metrics
The millis-since-base staleness check (last > now - 180s) saturated the
cutoff to 0 early in process life, so metrics stamped at 0ms failed the
strict > 0 test and were wrongly evicted by the immediate first GC
tick. Switch to age-based (now - last < 180s), equivalent to the
original Instant semantics and robust under clock saturation, which
fixes the flaky peer_conn_secure_mode_pubkey_and_encryption test.
2026-06-26 23:31:56 +08:00
fanyang 1375cd1832 docs(stats): fix misleading/stale comments and rename test_counter
- GC test: cutoff is in the future, so every metric is stale by timestamp
  (not "nothing is stale"); only live handles retain.
- GC loop: drop the inaccurate "no Instant alloc" rationale.
- bench: the handle path no longer calls Instant::now() (it uses fastant);
  reword the HANDLE_TOTAL_WORK rationale.
- rename test_unsafe_counter -> test_counter to match the type rename.
2026-06-25 09:16:12 +08:00
fanyang ebb97fd4f4 refactor(stats): rename UnsafeCounter to Counter
The counter is now a plain atomic (no longer UnsafeCell, no longer
sharded), so the "Unsafe" prefix is a misleading leftover. Rename to
Counter. Also fix a stale bench comment that claimed the ShardedAtomic
variant mirrors production.
2026-06-25 00:59:49 +08:00
fanyang 7d249979ea Update outdated comment 2026-06-25 00:54:26 +08:00
fanyang 8fa161bc1b Fix clippy 2026-06-25 00:51:55 +08:00
fanyang 2624b2740e Code format 2026-06-25 00:48:40 +08:00
fanyang bf8cff60bf docs: correct misleading multi_thread flag help text
The flag defaults to true and only affects launcher-based deployments
(GUI/mobile/web/Windows service); the easytier-core CLI is intentionally
single-threaded. Fix the help text and document the intent at the CLI entry.
2026-06-25 00:06:38 +08:00
fanyang a247358ec1 test(bench): add counter contention benchmark
Adds benches/counter_contention.rs comparing the production CounterHandle
against reconstructed baselines (pre-optimization single-atomic + CAS +
Mutex<Instant>) under tokio-task contention, plus counter-only variants.
Uses the real stats_manager types so the numbers reflect shipped code.

Placed between the sharding and fastant commits so each can be benchmarked
independently: at this commit prod = sharded counter + Mutex<Instant> touch.
2026-06-25 00:06:38 +08:00
fanyang c94d106714 perf(stats): lock-free fastant timestamp and fetch_add counter
Switch the counter add() from a fetch_update CAS loop to fetch_add, and
replace MetricData's Mutex<Instant> timestamp with a lock-free AtomicU64
storing millis since a lazily-initialized base. Back now_millis() with
fastant (TSC on x86_64 Linux, std fallback elsewhere) so touch() is cheap
enough to call per packet. GC and its test compare in the millis domain.

This is the change that actually delivers the performance: the Mutex
timestamp was the serialization bottleneck, and removing it (plus the TSC
clock) is what makes the handle path fast. No sharding.
2026-06-25 00:06:38 +08:00
fanyang 721b863547 fix: make stats counters thread safe 2026-06-25 00:06:38 +08:00
fanyangandGitHub 034f5066cd fix(faketcp): handle closed tun reader without panic (#2308)
Handle TUN receive errors by marking the fake TCP stack closed and
clearing registered sockets instead of panicking.

Refuse new sockets on closed stacks and let listeners recreate stacks
when the reader task exits.
2026-06-22 10:54:48 +08:00
fanyangandGitHub 9869ddaa4b fix: clarify config parse errors (#2360)
* fix: improve config parse diagnostics
* fix: polish config error context
* test: cover non-ascii config diagnostics
2026-06-21 21:56:49 +08:00
Luna YaoandGitHub 5ea6766238 fix: raise max_headers in ws handshake to 128 (#2366) 2026-06-21 21:54:07 +08:00
Luna YaoandGitHub 5efbc8587f upgrade guarden to 0.2.0 (#2365) 2026-06-18 23:45:28 +08:00
HYecandGitHub 7632cd64da Fix latency-first routing for direct peers (#2358) 2026-06-16 20:58:07 +08:00
韩嘉乐andGitHub 16b666ad25 fix: route_update message is not lag (#2355) 2026-06-16 00:00:48 +08:00
Luna YaoandGitHub 8909e88484 do not panic when fail to parse flags (#2349) 2026-06-14 13:12:20 +08:00
Luna YaoandGitHub 5edc4cb1cd fix: remove quinn-plaintext (#2345)
Remove quinn-plaintext to fix connection errors caused
by different hash values ​​across platforms.

On x64, maintain compatibility with quinn-plaintext.
2026-06-14 01:21:35 +08:00
韩嘉乐andGitHub e7709f1cb5 [OHOS] feat: improve status manager (#2343)
* fix: improve status manager
2026-06-11 17:27:10 +08:00
韩嘉乐andGitHub c0f42ebe8c fix: improve log manager (#2329) 2026-06-07 23:32:22 +08:00
KKRainbowandGitHub 9d965cae64 Add FFI JNI JSON RPC bridge (#2326)
* Add FFI JNI JSON RPC bridge
* Add FFI instance list API
2026-06-07 17:48:58 +08:00
韩嘉乐andGitHub da28c8badc [OHOS] fix: 修复内存泄露问题,并重构日志管理,预防性修复数据库初始化异常问题 (#2328)
* fix: leak memory
feat: new log manager

* fix: fail to init db

* fix: fail to init db

* fix: cargo format
2026-06-07 16:18:54 +08:00
KKRainbowandGitHub e38b1354b3 Fix credential ospf logic, fix udp subnet proxy loop protection (#2315) 2026-06-07 12:40:09 +08:00
KKRainbowandGitHub 793b57c2a1 feat(ffi): add async data plane API (#2321)
* feat(ffi): add async data plane API
* feat(ffi): add async data plane examples
* test(ffi): make async Go dataplane tests self-contained
* docs(ffi): document Go async dataplane API
* docs(android): document dataplane JNI API
2026-06-06 21:52:50 +08:00
深鸣andGitHub 4a25ca934b build: update pnpm config for v11 compatibility (#2322)
pnpm v11 introduces breaking changes that cause frozen installations
to fail:

1. The "pnpm" field in package.json is no longer read. Moved
   `overrides` to `pnpm-workspace.yaml` to fix
   `ERR_PNPM_LOCKFILE_CONFIG_MISMATCH`.
2. `strictDepBuilds` is now enabled by default. Added required
   dependencies (esbuild, unrs-resolver, vue-demi) to `allowBuilds` in
   the workspace config to fix `ERR_PNPM_IGNORED_BUILDS`.
2026-06-06 21:49:05 +08:00
KKRainbowandGitHub 13f2ebfe12 feat(ffi): add config server client bindings (#2320)
Add config server client support for the C FFI and Android JNI bindings.

Reuse the existing easytier::web_client::run_web_client path and 
NetworkInstanceManager; OHOS is unchanged.

Report successful remote config apply/delete operations through a 
callback, with one JSON event per affected instance.

Keep the config server client and FFI data plane mutually exclusive: once 
either side is in use, the other side returns an error instead of sharing
lifecycle state.
2026-06-06 01:52:27 +08:00
Luna YaoandGitHub 9ba364ff60 feat: string deserialization for prost enums (#2316)
Use pbjson to support string deserialization for enum fields

This allows TOML configs like:
    chainType = "Inbound"
instead of:
    chainType = 1

- Maintain backward compatibility with integer values
- Default serialization format is now string
2026-06-06 00:21:22 +08:00
NeilandGitHub ba653da9a0 fix: detect credential mode in TOML config loader (#2301)
TomlConfigLoader::new_from_str() always calls NetworkIdentity::new()
with unwrap_or_default() on network_secret, converting None to ''.
This creates a non-zero SHA256 digest, causing credential nodes loaded
from TOML to be misidentified as regular nodes (with network_secret),
which breaks Noise handshake authentication.

Fix: check if secure_mode is enabled AND network_secret is absent/empty,
and call NetworkIdentity::new_credential() in that case.

The same detection already exists in:
- core.rs (CLI path, via --credential flag)
- launcher.rs (GUI/web path, via gen_config)

This makes TOML config loading consistent with the other two entry points.
2026-06-04 22:38:40 +08:00
Luna YaoandGitHub 64c4d73044 fix QuicSocket payload offset (#2306) 2026-06-04 18:04:39 +08:00
w568wandGitHub e0745f4bab feat: Add data plane support to FFI (#2287)
1. Overview

This PR adds data plane APIs to easytier-ffi:

TCP Outbound:

- data_plane_tcp_connect
- data_plane_tcp_read
- data_plane_tcp_write
- data_plane_tcp_close

TCP Listener:

- data_plane_tcp_bind
- data_plane_tcp_accept
- data_plane_tcp_listener_close

UDP:

- data_plane_udp_bind
- data_plane_udp_send_to
- data_plane_udp_recv_from
- data_plane_udp_close

2. Key Changes

The main changes are focused on:

- easytier-contrib/easytier-ffi/src/lib.rs: Added FFI interfaces;
  made ERROR_MSG thread-safe.

- easytier/src/gateway/socks5.rs: Bridges the data plane to the
  existing Socks5 server logic.
  - Added EasyTierUdpSocket, mainly wrapping ref-counting and
    critical object (e.g., Socks5EntrySet) hold & drop logic,
    and exposing common fields (e.g., local_addr).
  - Extended Socks5Server functionality to expose TCP and UDP
    socket creation interfaces for FFI calls.

- Other files: Mostly pass-through logic.

- Added a relatively large Go usage example.
2026-06-04 17:17:41 +08:00
e3ca7ffa54 feat(socket): add Linux SO_MARK (fwmark) support for underlay sockets (#2288)
Adds a Linux-only socket_mark u32 config flag (CLI: --socket-mark, env:
ET_SOCKET_MARK, TOML/proto: flags.socket_mark, 0 = disabled) that is
applied as SO_MARK to every outbound underlay socket EasyTier creates:
TCP, UDP, QUIC, WebSocket, WireGuard connectors and listeners, plus the
FakeTCP decoy socket. Lets the host policy-route or filter EasyTier
underlay traffic with 'ip rule fwmark ...' or iptables -m mark.

Plumbing mirrors the existing bind_device pattern:
- FlagsInConfig.socket_mark (proto) + default 0 in gen_default_flags
- bind() builder gets a socket_mark arg; setup_socket2_ext calls
  apply_socket_mark which is a no-op for mark=0 and on non-Linux
- TunnelConnector trait gets set_socket_mark(u32) default-no-op method
- IP-based connectors override; create_listener_by_url and the connector
  factory pass mark from global_ctx flags
- QUIC threads mark through QuicEndpointManager::{server,connect}
- WebSocket/FakeTCP/TCP default-bind bypass paths apply mark via
  socket2::SockRef::from(&tokio_socket)
- ForeignNetworkEntry propagates parent socket_mark into its derived ctx

Includes a Linux smoke test plus a CAP_NET_ADMIN-gated test that does a
getsockopt(SO_MARK) round-trip to confirm the kernel applied the value.

SO_MARK requires CAP_NET_ADMIN; ignored silently on non-Linux. FakeTCP's
TUN-written segments are not covered (kernel doesn't tag raw TUN
writes); operators relying on fwmark for FakeTCP must apply an iptables
rule on the FakeTCP TUN device separately.

Co-authored-by: Claude <noreply@anthropic.com>
2026-06-03 09:57:18 +08:00
韩嘉乐andGitHub df97f3a64d fix: make ohos snapshot sync and config id validation safer (#2283)
* feat: add the management of config_store_snapshot

* fix: make ohrs snapshot sync and config id validation safer
2026-06-02 22:40:48 +08:00
fanyangandGitHub 00957e5f9d feat: add Docker healthcheck (#2279) 2026-05-24 23:12:51 +08:00
fanyangandGitHub bfa3383aaa chore: update kcp-sys (#2277) 2026-05-24 23:11:16 +08:00
121 changed files with 18616 additions and 2240 deletions
+3
View File
@@ -42,4 +42,7 @@ EXPOSE 11011/tcp
# wss
EXPOSE 11012/tcp
HEALTHCHECK --interval=30s --timeout=10s --start-period=60s --retries=5 \
CMD ["/usr/local/bin/easytier-cli", "--rpc-portal", "127.0.0.1:15888", "--output", "json", "node", "info"]
ENTRYPOINT ["/sbin/tini", "--", "easytier-core"]
+3
View File
@@ -43,3 +43,6 @@ easytier-gui/src-tauri/*.sys
.direnv
.flake-profile
# contrib
go.sum
Generated
+327 -90
View File
@@ -129,6 +129,12 @@ dependencies = [
"libc",
]
[[package]]
name = "anes"
version = "0.1.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299"
[[package]]
name = "anstream"
version = "0.6.15"
@@ -241,6 +247,16 @@ dependencies = [
"password-hash",
]
[[package]]
name = "ariadne"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "36f5e3dca4e09a6f340a61a0e9c7b61e030c69fc27bf29d73218f7e5e3b7638f"
dependencies = [
"unicode-width 0.1.11",
"yansi",
]
[[package]]
name = "arrayvec"
version = "0.7.6"
@@ -915,7 +931,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bfcfdc083699101d5a7965e49925975f2f55060f94f9a05e7187be95d530ca59"
dependencies = [
"once_cell",
"proc-macro-crate 3.2.0",
"proc-macro-crate 3.5.0",
"proc-macro2",
"quote",
"syn 2.0.117",
@@ -1129,6 +1145,12 @@ dependencies = [
"toml 0.9.12+spec-1.1.0",
]
[[package]]
name = "cast"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5"
[[package]]
name = "cc"
version = "1.2.10"
@@ -1238,6 +1260,33 @@ dependencies = [
"windows-targets 0.52.6",
]
[[package]]
name = "ciborium"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e"
dependencies = [
"ciborium-io",
"ciborium-ll",
"serde",
]
[[package]]
name = "ciborium-io"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757"
[[package]]
name = "ciborium-ll"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9"
dependencies = [
"ciborium-io",
"half",
]
[[package]]
name = "cidr"
version = "0.3.1"
@@ -1407,8 +1456,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8030735ecb0d128428b64cd379809817e620a40e5001c54465b99ec5feec2857"
dependencies = [
"futures-core",
"prost",
"prost-types",
"prost 0.13.5",
"prost-types 0.13.5",
"tonic",
"tracing-core",
]
@@ -1426,8 +1475,8 @@ dependencies = [
"hdrhistogram",
"humantime",
"hyper-util",
"prost",
"prost-types",
"prost 0.13.5",
"prost-types 0.13.5",
"serde",
"serde_json",
"thread_local",
@@ -1583,6 +1632,42 @@ dependencies = [
"cfg-if",
]
[[package]]
name = "criterion"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f2b12d017a929603d80db1831cd3a24082f8137ce19c69e6447f54f5fc8d692f"
dependencies = [
"anes",
"cast",
"ciborium",
"clap",
"criterion-plot",
"is-terminal",
"itertools 0.10.5",
"num-traits",
"once_cell",
"oorandom",
"plotters",
"rayon",
"regex",
"serde",
"serde_derive",
"serde_json",
"tinytemplate",
"walkdir",
]
[[package]]
name = "criterion-plot"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6b50826342786a51a89e2da3a28f1c32b06e387201bc2d19791f622c673706b1"
dependencies = [
"cast",
"itertools 0.10.5",
]
[[package]]
name = "critical-section"
version = "1.2.0"
@@ -2234,6 +2319,7 @@ dependencies = [
"aes-gcm",
"anyhow",
"arc-swap",
"ariadne",
"async-recursion",
"async-ringbuf",
"async-stream",
@@ -2255,6 +2341,7 @@ dependencies = [
"clap_complete",
"clap_complete_nushell",
"console-subscriber",
"criterion",
"crossbeam",
"ctor 0.8.0",
"dashmap",
@@ -2265,6 +2352,7 @@ dependencies = [
"derive_builder",
"derive_more 2.1.1",
"encoding",
"fastant",
"flume 0.12.0",
"forwarded-header-value",
"futures",
@@ -2272,7 +2360,7 @@ dependencies = [
"gethostname 0.5.0",
"git-version",
"globwalk",
"guarden",
"guarden 0.2.0",
"hickory-client",
"hickory-proto",
"hickory-resolver",
@@ -2304,21 +2392,21 @@ dependencies = [
"ordered_hash_map",
"parking_lot",
"paste",
"pbjson",
"pbjson-build",
"percent-encoding",
"petgraph 0.8.1",
"petgraph",
"pin-project-lite",
"pnet",
"prefix-trie",
"proc-macro2",
"prost",
"prost 0.14.3",
"prost-build",
"prost-reflect",
"prost-reflect-build",
"prost-wkt",
"prost-wkt-build",
"prost-wkt-types",
"quinn",
"quinn-plaintext",
"quinn-proto",
"quote",
"rand 0.8.5",
"rcgen",
@@ -2330,6 +2418,7 @@ dependencies = [
"rstest",
"rust-i18n",
"rustls",
"seahash",
"serde",
"serde_json",
"serial_test",
@@ -2385,6 +2474,7 @@ version = "0.1.0"
dependencies = [
"android_logger",
"easytier",
"easytier-ffi",
"jni",
"log",
"once_cell",
@@ -2396,11 +2486,17 @@ dependencies = [
name = "easytier-ffi"
version = "0.1.0"
dependencies = [
"async-trait",
"dashmap",
"easytier",
"log",
"once_cell",
"percent-encoding",
"serde",
"serde_json",
"tokio",
"tokio-util",
"url",
"uuid",
]
@@ -2450,7 +2546,7 @@ dependencies = [
"dashmap",
"easytier",
"futures",
"guarden",
"guarden 0.1.2",
"jsonwebtoken",
"mimalloc",
"mockall",
@@ -2829,6 +2925,16 @@ dependencies = [
"pin-project-lite",
]
[[package]]
name = "fastant"
version = "0.1.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2e825441bfb2d831c47c97d05821552db8832479f44c571b97fededbf0099c07"
dependencies = [
"small_ctor",
"web-time",
]
[[package]]
name = "fastbloom"
version = "0.9.0"
@@ -2905,12 +3011,6 @@ dependencies = [
"rustc_version",
]
[[package]]
name = "fixedbitset"
version = "0.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0ce7134b9999ecaf8bcd65542e436736ef32ddca1b3e06094cb6ec5755203b80"
[[package]]
name = "fixedbitset"
version = "0.5.7"
@@ -3592,7 +3692,18 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ca87812d87fa82896df1adfb5c111cdeaae3edb6da028f5df002dcbd7df71454"
dependencies = [
"futures",
"guarden-macros",
"guarden-macros 0.1.2",
"tokio",
]
[[package]]
name = "guarden"
version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b8408903291a7d0cc74169d5de4dd1919a9a402a2f67fcd7df3303ed045fae73"
dependencies = [
"futures-core",
"guarden-macros 0.2.0",
"tokio",
]
@@ -3607,6 +3718,18 @@ dependencies = [
"syn 2.0.117",
]
[[package]]
name = "guarden-macros"
version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1e0ef28f1077c259f9e7e238e234a78ce18cedbf0251fd2135f5fc23c40e79fe"
dependencies = [
"proc-macro-crate 3.5.0",
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "h2"
version = "0.4.7"
@@ -4037,7 +4160,7 @@ dependencies = [
"libc",
"percent-encoding",
"pin-project-lite",
"socket2 0.5.10",
"socket2 0.6.3",
"tokio",
"tower-service",
"tracing",
@@ -4381,9 +4504,9 @@ dependencies = [
[[package]]
name = "inventory"
version = "0.3.22"
version = "0.3.24"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "009ae045c87e7082cb72dab0ccd01ae075dd00141ddc108f43a0ea150a9e7227"
checksum = "a4f0c30c76f2f4ccee3fe55a2435f691ca00c0e4bd87abe4f4a851b1d4dac39b"
dependencies = [
"rustversion",
]
@@ -4459,6 +4582,17 @@ dependencies = [
"once_cell",
]
[[package]]
name = "is-terminal"
version = "0.4.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46"
dependencies = [
"hermit-abi",
"libc",
"windows-sys 0.61.2",
]
[[package]]
name = "is-wsl"
version = "0.4.0"
@@ -4623,7 +4757,7 @@ dependencies = [
[[package]]
name = "kcp-sys"
version = "0.1.0"
source = "git+https://github.com/EasyTier/kcp-sys?rev=94964794caaed5d388463137da59b97499619e5f#94964794caaed5d388463137da59b97499619e5f"
source = "git+https://github.com/EasyTier/kcp-sys?rev=d7427c22d764deb1860a7d37acc446ed5033464c#d7427c22d764deb1860a7d37acc446ed5033464c"
dependencies = [
"anyhow",
"auto_impl",
@@ -5577,7 +5711,7 @@ version = "0.7.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "680998035259dcfcafe653688bf2aa6d3e2dc05e98be6ab46afb089dc84f1df8"
dependencies = [
"proc-macro-crate 3.2.0",
"proc-macro-crate 3.5.0",
"proc-macro2",
"quote",
"syn 2.0.117",
@@ -5819,6 +5953,12 @@ dependencies = [
"portable-atomic",
]
[[package]]
name = "oorandom"
version = "11.1.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e"
[[package]]
name = "opaque-debug"
version = "0.3.1"
@@ -6154,6 +6294,28 @@ version = "0.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df94ce210e5bc13cb6651479fa48d14f601d9858cfe0467f43ae157023b938d3"
[[package]]
name = "pbjson"
version = "0.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e8edd1efdd8ab23ba9cb9ace3d9987a72663d5d7c9f74fa00b51d6213645cf6c"
dependencies = [
"base64 0.22.1",
"serde",
]
[[package]]
name = "pbjson-build"
version = "0.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2ed4d5c6ae95e08ac768883c8401cf0e8deb4e6e1d6a4e1fd3d2ec4f0ec63200"
dependencies = [
"heck 0.5.0",
"itertools 0.14.0",
"prost 0.14.3",
"prost-types 0.14.3",
]
[[package]]
name = "pbkdf2"
version = "0.12.2"
@@ -6189,23 +6351,13 @@ version = "2.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e3148f5046208a5d56bcfc03053e3ca6334e51da8dfb19b6cdc8b306fae3283e"
[[package]]
name = "petgraph"
version = "0.6.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b4c5cc86750666a3ed20bdaf5ca2a0344f9c67674cae0515bec2da16fbaa47db"
dependencies = [
"fixedbitset 0.4.2",
"indexmap 2.14.0",
]
[[package]]
name = "petgraph"
version = "0.8.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7a98c6720655620a521dcc722d0ad66cd8afd5d86e34a89ef691c50b7b24de06"
dependencies = [
"fixedbitset 0.5.7",
"fixedbitset",
"hashbrown 0.15.3",
"indexmap 2.14.0",
"serde",
@@ -6437,6 +6589,34 @@ dependencies = [
"time",
]
[[package]]
name = "plotters"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747"
dependencies = [
"num-traits",
"plotters-backend",
"plotters-svg",
"wasm-bindgen",
"web-sys",
]
[[package]]
name = "plotters-backend"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a"
[[package]]
name = "plotters-svg"
version = "0.3.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670"
dependencies = [
"plotters-backend",
]
[[package]]
name = "pnet"
version = "0.35.0"
@@ -6697,11 +6877,11 @@ dependencies = [
[[package]]
name = "proc-macro-crate"
version = "3.2.0"
version = "3.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8ecf48c7ca261d60b74ab1a7b20da18bede46776b2e55535cb958eb595c5fa7b"
checksum = "e67ba7e9b2b56446f1d419b1d807906278ffa1a658a8a5d8a39dcb1f5a78614f"
dependencies = [
"toml_edit 0.22.20",
"toml_edit 0.25.12+spec-1.1.0",
]
[[package]]
@@ -6785,24 +6965,33 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2796faa41db3ec313a31f7624d9286acf277b52de526150b7e69f3debf891ee5"
dependencies = [
"bytes",
"prost-derive",
"prost-derive 0.13.5",
]
[[package]]
name = "prost"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d2ea70524a2f82d518bce41317d0fae74151505651af45faf1ffbd6fd33f0568"
dependencies = [
"bytes",
"prost-derive 0.14.3",
]
[[package]]
name = "prost-build"
version = "0.13.5"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be769465445e8c1474e9c5dac2018218498557af32d9ed057325ec9a41ae81bf"
checksum = "343d3bd7056eda839b03204e68deff7d1b13aba7af2b2fd16890697274262ee7"
dependencies = [
"heck 0.5.0",
"itertools 0.14.0",
"log",
"multimap",
"once_cell",
"petgraph 0.6.5",
"petgraph",
"prettyplease",
"prost",
"prost-types",
"prost 0.14.3",
"prost-types 0.14.3",
"regex",
"syn 2.0.117",
"tempfile",
@@ -6822,22 +7011,34 @@ dependencies = [
]
[[package]]
name = "prost-reflect"
version = "0.14.5"
name = "prost-derive"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e92b959d24e05a3e2da1d0beb55b48bc8a97059b8336ea617780bd6addbbfb5a"
checksum = "27c6023962132f4b30eb4c172c91ce92d933da334c59c23cddee82358ddafb0b"
dependencies = [
"once_cell",
"prost",
"anyhow",
"itertools 0.14.0",
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "prost-reflect"
version = "0.16.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "590aa145fee8f7a26b5a6055365e7c5e89a5c1caae9869de76ec0ee73181a2f9"
dependencies = [
"prost 0.14.3",
"prost-reflect-derive",
"prost-types",
"prost-types 0.14.3",
]
[[package]]
name = "prost-reflect-build"
version = "0.14.0"
version = "0.16.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "50e2537231d94dd2778920c2ada37dd9eb1ac0325bb3ee3ee651bd44c1134123"
checksum = "8214ae2c30bbac390db0134d08300e770ef89b6d4e5abf855e8d300eded87e28"
dependencies = [
"prost-build",
"prost-reflect",
@@ -6845,9 +7046,9 @@ dependencies = [
[[package]]
name = "prost-reflect-derive"
version = "0.14.0"
version = "0.16.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f4fce6b22f15cc8d8d400a2b98ad29202b33bd56c7d9ddd815bc803a807ecb65"
checksum = "7b6d90e29fa6c0d13c2c19ba5e4b3fb0efbf5975d27bcf4e260b7b15455bcabe"
dependencies = [
"proc-macro2",
"quote",
@@ -6860,18 +7061,27 @@ version = "0.13.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "52c2c1bf36ddb1a1c396b3601a3cec27c2462e45f07c386894ec3ccf5332bd16"
dependencies = [
"prost",
"prost 0.13.5",
]
[[package]]
name = "prost-types"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8991c4cbdb8bc5b11f0b074ffe286c30e523de90fee5ba8132f1399f23cb3dd7"
dependencies = [
"prost 0.14.3",
]
[[package]]
name = "prost-wkt"
version = "0.6.1"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "497e1e938f0c09ef9cabe1d49437b4016e03e8f82fbbe5d1c62a9b61b9decae1"
checksum = "cd3de5e9c9e84fcb5efa204b8e283d23e615a8bc8c777bf1d6622bb01dc61445"
dependencies = [
"chrono",
"inventory",
"prost",
"prost 0.14.3",
"serde",
"serde_derive",
"serde_json",
@@ -6880,27 +7090,27 @@ dependencies = [
[[package]]
name = "prost-wkt-build"
version = "0.6.1"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "07b8bf115b70a7aa5af1fd5d6e9418492e9ccb6e4785e858c938e28d132a884b"
checksum = "fe500dc80e757a75e1e8fb7290e448d62dfba3105ece1d058579cb00b58151cd"
dependencies = [
"heck 0.5.0",
"prost",
"prost 0.14.3",
"prost-build",
"prost-types",
"prost-types 0.14.3",
"quote",
]
[[package]]
name = "prost-wkt-types"
version = "0.6.1"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c8cdde6df0a98311c839392ca2f2f0bcecd545f86a62b4e3c6a49c336e970fe5"
checksum = "13807eaa7e15833d06e899008371926201cdcd11d74b6d490f49130cdb3f415e"
dependencies = [
"chrono",
"prost",
"prost 0.14.3",
"prost-build",
"prost-types",
"prost-types 0.14.3",
"prost-wkt",
"prost-wkt-build",
"regex",
@@ -6979,18 +7189,6 @@ dependencies = [
"web-time",
]
[[package]]
name = "quinn-plaintext"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f3e617feaeb6493018fa35fc47ae8b630ac8903d8159e9e747018841b99bad3d"
dependencies = [
"bytes",
"quinn-proto",
"seahash",
"tracing",
]
[[package]]
name = "quinn-proto"
version = "0.11.12"
@@ -7565,7 +7763,7 @@ checksum = "1f168d99749d307be9de54d23fd226628d99768225ef08f6ffb52e0182a27746"
dependencies = [
"cfg-if",
"glob",
"proc-macro-crate 3.2.0",
"proc-macro-crate 3.5.0",
"proc-macro2",
"quote",
"regex",
@@ -8612,6 +8810,12 @@ dependencies = [
"autocfg",
]
[[package]]
name = "small_ctor"
version = "0.1.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "88414a5ca1f85d82cc34471e975f0f74f6aa54c40f062efa42c0080e7f763f81"
[[package]]
name = "smallvec"
version = "1.13.2"
@@ -9770,6 +9974,16 @@ dependencies = [
"zerovec",
]
[[package]]
name = "tinytemplate"
version = "1.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc"
dependencies = [
"serde",
"serde_json",
]
[[package]]
name = "tinyvec"
version = "1.8.0"
@@ -9889,8 +10103,7 @@ dependencies = [
[[package]]
name = "tokio-websockets"
version = "0.13.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dad543404f98bfc969aeb71994105c592acfc6c43323fddcd016bb208d1c65cb"
source = "git+https://github.com/EasyTier/tokio-websockets#dc9771c7c215882349c3cb328877550a3593df21"
dependencies = [
"base64 0.22.1",
"bytes",
@@ -9965,6 +10178,15 @@ dependencies = [
"serde_core",
]
[[package]]
name = "toml_datetime"
version = "1.1.1+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3165f65f62e28e0115a00b2ebdd37eb6f3b641855f9d636d3cd4103767159ad7"
dependencies = [
"serde_core",
]
[[package]]
name = "toml_edit"
version = "0.19.15"
@@ -10002,6 +10224,18 @@ dependencies = [
"winnow 0.6.18",
]
[[package]]
name = "toml_edit"
version = "0.25.12+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d2153edc6955a6c354fad8f5efd38b6a8769bdccf9fe50f8e1329f81b0baa5d7"
dependencies = [
"indexmap 2.14.0",
"toml_datetime 1.1.1+spec-1.1.0",
"toml_parser",
"winnow 1.0.1",
]
[[package]]
name = "toml_parser"
version = "1.1.2+spec-1.1.0"
@@ -10037,7 +10271,7 @@ dependencies = [
"hyper-util",
"percent-encoding",
"pin-project",
"prost",
"prost 0.13.5",
"socket2 0.5.10",
"tokio",
"tokio-stream",
@@ -10321,7 +10555,7 @@ checksum = "b8765b90061cba6c22b5831f675da109ae5561588290f9fa2317adab2714d5a6"
dependencies = [
"memchr",
"nom 8.0.0",
"petgraph 0.8.1",
"petgraph",
]
[[package]]
@@ -10381,9 +10615,9 @@ checksum = "42ff0bf0c66b8238c6f3b578df37d0b7848e55df8577b3f74f92a69acceeb825"
[[package]]
name = "typetag"
version = "0.2.21"
version = "0.2.22"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be2212c8a9b9bcfca32024de14998494cf9a5dfa59ea1b829de98bac374b86bf"
checksum = "c5a897b12c6c1151ad0b138b8db50252dc301f93bc3b027db05eec82aeed298c"
dependencies = [
"erased-serde",
"inventory",
@@ -10394,9 +10628,9 @@ dependencies = [
[[package]]
name = "typetag-impl"
version = "0.2.21"
version = "0.2.22"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "27a7a9b72ba121f6f1f6c3632b85604cac41aedb5ddc70accbebb6cac83de846"
checksum = "cf808357c6ed7e13ba0f3277ec8d8f21b2d501274895104263985330c726c1c5"
dependencies = [
"proc-macro2",
"quote",
@@ -11849,6 +12083,9 @@ name = "winnow"
version = "1.0.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "09dac053f1cd375980747450bfc7250c264eaae0583872e845c0c7cd578872b5"
dependencies = [
"memchr",
]
[[package]]
name = "winreg"
@@ -12230,7 +12467,7 @@ version = "5.14.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "897e79616e84aac4b2c46e9132a4f63b93105d54fe8c0e8f6bffc21fa8d49222"
dependencies = [
"proc-macro-crate 3.2.0",
"proc-macro-crate 3.5.0",
"proc-macro2",
"quote",
"syn 2.0.117",
@@ -12467,7 +12704,7 @@ version = "5.10.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5b59b012ebe9c46656f9cc08d8da8b4c726510aef12559da3e5f1bf72780752c"
dependencies = [
"proc-macro-crate 3.2.0",
"proc-macro-crate 3.5.0",
"proc-macro2",
"quote",
"syn 2.0.117",
@@ -13,4 +13,5 @@ log = "0.4"
android_logger = "0.13"
serde = { version = "1.0.220", features = ["derive"] }
serde_json = "1.0"
easytier = { path = "../../easytier" }
easytier = { path = "../../easytier" }
easytier-ffi = { path = "../easytier-ffi", default-features = false, features = ["ffi-dataplane"] }
@@ -8,6 +8,7 @@
- 📱 原生 Android JNI 支持
- 🔧 支持多种 Android 架构 (arm64-v8a, armeabi-v7a, x86, x86_64)
- 🛡️ 类型安全的 Java 接口
- 🔌 支持通过 JSON 调用已暴露的 EasyTier RPC 查询/管理接口
- 📝 详细的错误处理和日志记录
## 支持的架构
@@ -176,6 +177,20 @@ public class EasyTierManager {
}
```
### 通用 JSON RPC
`EasyTierJNI.callJsonRpc(serviceName, methodName, domainName, payloadJson)` 可以调用已暴露的
EasyTier RPC 服务,payload 和返回值均为 protobuf JSON。该接口不支持
`api.manage.WebClientService`;实例启动、保留、删除、信息收集仍使用专用 JNI API。
```java
String response = EasyTierJNI.callJsonRpc(
"api.logger.LoggerRpcService",
"get_logger_config",
"{}"
);
```
### VPN 服务集成
如果您要在 Android VPN 服务中使用:
@@ -264,4 +279,4 @@ public class EasyTierVpnService extends VpnService {
- [EasyTier 主项目](https://github.com/EasyTier/EasyTier)
- [Android NDK 文档](https://developer.android.com/ndk)
- [Rust JNI 文档](https://docs.rs/jni/)
- [Rust JNI 文档](https://docs.rs/jni/)
@@ -0,0 +1,17 @@
use std::{env, path::PathBuf};
fn main() {
let target_os = env::var("CARGO_CFG_TARGET_OS").unwrap_or_default();
if !matches!(target_os.as_str(), "android" | "linux") {
return;
}
let manifest_dir = PathBuf::from(env::var("CARGO_MANIFEST_DIR").unwrap());
let exports = manifest_dir.join("exports.map");
println!("cargo:rerun-if-changed={}", exports.display());
println!(
"cargo:rustc-cdylib-link-arg=-Wl,--version-script={}",
exports.display()
);
println!("cargo:rustc-cdylib-link-arg=-Wl,--exclude-libs,ALL");
}
@@ -0,0 +1,7 @@
{
global:
Java_com_easytier_jni_EasyTierJNI_*;
Java_com_easytier_jni_EasyTierDataPlaneJNI_*;
local:
*;
};
@@ -0,0 +1,451 @@
package com.easytier.jni
import kotlinx.coroutines.CancellationException
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.currentCoroutineContext
import kotlinx.coroutines.ensureActive
import kotlinx.coroutines.withContext
/**
* EasyTier data-plane API for Android.
*
* Dataplane APIs do not create or start an EasyTier instance by themselves.
* Start an instance with [EasyTierJNI.runNetworkInstance] first, then pass the
* same `instanceName` to [EasyTierDataPlane.tcpConnect],
* [EasyTierDataPlane.tcpBind], or [EasyTierDataPlane.udpBind]. If that instance
* is not running, the native start call fails and the coroutine wrapper throws
* the last EasyTier FFI error.
*
* Typical setup:
* ```
* val instanceName = "android-dataplane-demo"
* val config = """
* instance_name = "$instanceName"
* ipv4 = "10.144.0.1"
* listeners = ["tcp://0.0.0.0:11010"]
*
* [network_identity]
* network_name = "android-dataplane-demo"
* network_secret = "replace-with-a-real-secret"
*
* [[peer]]
* uri = "tcp://peer.example.com:11010"
*
* [flags]
* no_tun = true
* bind_device = false
* """.trimIndent()
*
* EasyTierJNI.runNetworkInstance(config)
* ```
*
* After the instance is running, most callers should use [EasyTierDataPlane]
* and the socket/stream classes below. [EasyTierDataPlaneJNI] is the low-level
* native op-handle ABI used by the coroutine wrappers.
*
* TCP client usage:
* ```
* val stream = EasyTierDataPlane.tcpConnect(instanceName, "10.144.0.2", 8080, 5_000)
* try {
* stream.write("ping".toByteArray(), 5_000)
* val reply = stream.read(4096, 5_000)
* } finally {
* stream.close()
* }
* ```
*
* TCP server usage:
* ```
* val listener = EasyTierDataPlane.tcpBind(instanceName, 8080, 5_000)
* try {
* val stream = listener.accept(30_000)
* try {
* stream.write(stream.read(4096, 5_000), 5_000)
* } finally {
* stream.close()
* }
* } finally {
* listener.close()
* }
* ```
*
* UDP usage:
* ```
* val socket = EasyTierDataPlane.udpBind(instanceName, 0, 5_000)
* try {
* socket.sendTo("10.144.0.2", 9000, "ping".toByteArray(), 5_000)
* val packet = socket.recvFrom(4096, 5_000)
* } finally {
* socket.close()
* }
* ```
*
* Operation model:
* - Each suspend function starts one native async op, waits on Dispatchers.IO,
* then consumes the op with the matching finish call.
* - Coroutine cancellation cancels and frees the native op.
* - Returned stream/listener/socket handles must be closed by the caller.
* - Input ByteArray data is copied by the native start call; output data is
* copied into Kotlin ByteArray before the native buffer is freed.
*/
/** Data-plane IPv4/port pair returned by EasyTier FFI. */
data class DataPlaneSocketAddress(val ip: String, val port: Int)
/** Result of a completed TCP connect op. */
data class DataPlaneTcpConnectResult(val handle: Long, val localAddress: DataPlaneSocketAddress)
/** Result of a completed TCP bind op. */
data class DataPlaneTcpBindResult(val handle: Long, val localAddress: DataPlaneSocketAddress)
/** Result of a completed TCP accept op. */
data class DataPlaneTcpAcceptResult(
val handle: Long,
val localAddress: DataPlaneSocketAddress,
val peerAddress: DataPlaneSocketAddress
)
/** Result of a completed TCP read op. */
data class DataPlaneTcpReadResult(val data: ByteArray)
/** Result of a completed UDP bind op. */
data class DataPlaneUdpBindResult(val handle: Long, val localAddress: DataPlaneSocketAddress)
/** Result of a completed UDP recv_from op. */
data class DataPlaneUdpRecvResult(
val data: ByteArray,
val peerAddress: DataPlaneSocketAddress
)
/** TCP data-plane stream handle. Call [close] when the stream is no longer needed. */
class DataPlaneTcpStream(
val handle: Long,
val localAddress: DataPlaneSocketAddress? = null,
val peerAddress: DataPlaneSocketAddress? = null
) {
/** Read up to [maxLength] bytes, waiting at most [timeoutMs] in native code. */
suspend fun read(maxLength: Int, timeoutMs: Long): ByteArray =
EasyTierDataPlane.tcpRead(this, maxLength, timeoutMs)
/** Write [data], waiting at most [timeoutMs] in native code. */
suspend fun write(data: ByteArray, timeoutMs: Long): Int =
EasyTierDataPlane.tcpWrite(this, data, timeoutMs)
/** Close the native TCP stream handle. */
fun close(): Int = EasyTierDataPlaneJNI.dataPlaneTcpClose(handle)
}
/** TCP data-plane listener handle. Call [close] when the listener is no longer needed. */
class DataPlaneTcpListener(val handle: Long, val localAddress: DataPlaneSocketAddress) {
/** Accept one TCP data-plane stream. */
suspend fun accept(timeoutMs: Long): DataPlaneTcpStream =
EasyTierDataPlane.tcpAccept(this, timeoutMs)
/** Close the native TCP listener handle. */
fun close(): Int = EasyTierDataPlaneJNI.dataPlaneTcpListenerClose(handle)
}
/** UDP data-plane socket handle. Call [close] when the socket is no longer needed. */
class DataPlaneUdpSocket(val handle: Long, val localAddress: DataPlaneSocketAddress) {
/** Send one UDP datagram to [dstIp]:[dstPort]. */
suspend fun sendTo(
dstIp: String,
dstPort: Int,
data: ByteArray,
timeoutMs: Long
): Int = EasyTierDataPlane.udpSendTo(this, dstIp, dstPort, data, timeoutMs)
/** Receive one UDP datagram and its peer address. */
suspend fun recvFrom(maxLength: Int, timeoutMs: Long): DataPlaneUdpRecvResult =
EasyTierDataPlane.udpRecvFrom(this, maxLength, timeoutMs)
/** Close the native UDP socket handle. */
fun close(): Int = EasyTierDataPlaneJNI.dataPlaneUdpClose(handle)
}
/**
* Low-level native data-plane JNI entry points.
*
* These functions mirror the Rust FFI op-handle ABI directly. They are exposed
* for completeness, but most Android callers should use [EasyTierDataPlane]
* instead so coroutine cancellation and op cleanup are handled consistently.
*/
object EasyTierDataPlaneJNI {
init {
System.loadLibrary("easytier_android_jni")
}
@JvmStatic external fun dataPlaneAsyncOpStatus(handle: Long): Int
@JvmStatic external fun dataPlaneAsyncOpWait(handle: Long, timeoutMs: Long): Int
@JvmStatic external fun dataPlaneAsyncOpCancel(handle: Long): Int
@JvmStatic external fun dataPlaneAsyncOpFree(handle: Long): Int
@JvmStatic
external fun dataPlaneTcpConnectStart(
instanceName: String,
dstIp: String,
dstPort: Int,
timeoutMs: Long
): Long
@JvmStatic external fun dataPlaneTcpConnectFinish(op: Long): DataPlaneTcpConnectResult?
@JvmStatic
external fun dataPlaneTcpBindStart(
instanceName: String,
localPort: Int,
timeoutMs: Long
): Long
@JvmStatic external fun dataPlaneTcpBindFinish(op: Long): DataPlaneTcpBindResult?
@JvmStatic external fun dataPlaneTcpAcceptStart(handle: Long, timeoutMs: Long): Long
@JvmStatic external fun dataPlaneTcpAcceptFinish(op: Long): DataPlaneTcpAcceptResult?
@JvmStatic external fun dataPlaneTcpReadStart(handle: Long, maxLength: Int, timeoutMs: Long): Long
@JvmStatic external fun dataPlaneTcpReadFinish(op: Long): DataPlaneTcpReadResult?
@JvmStatic external fun dataPlaneTcpWriteStart(handle: Long, data: ByteArray, timeoutMs: Long): Long
@JvmStatic external fun dataPlaneTcpWriteFinish(op: Long): Int
@JvmStatic
external fun dataPlaneUdpBindStart(
instanceName: String,
localPort: Int,
timeoutMs: Long
): Long
@JvmStatic external fun dataPlaneUdpBindFinish(op: Long): DataPlaneUdpBindResult?
@JvmStatic
external fun dataPlaneUdpSendToStart(
handle: Long,
dstIp: String,
dstPort: Int,
data: ByteArray,
timeoutMs: Long
): Long
@JvmStatic external fun dataPlaneUdpSendToFinish(op: Long): Int
@JvmStatic external fun dataPlaneUdpRecvFromStart(handle: Long, maxLength: Int, timeoutMs: Long): Long
@JvmStatic external fun dataPlaneUdpRecvFromFinish(op: Long): DataPlaneUdpRecvResult?
@JvmStatic external fun dataPlaneTcpClose(handle: Long): Int
@JvmStatic external fun dataPlaneTcpListenerClose(handle: Long): Int
@JvmStatic external fun dataPlaneUdpClose(handle: Long): Int
}
/** Coroutine-friendly Android data-plane API. */
object EasyTierDataPlane {
private const val DATA_PLANE_OP_PENDING = 0
private const val DATA_PLANE_OP_READY = 1
private const val DATA_PLANE_OP_FAILED = -1
private const val DATA_PLANE_OP_INVALID = -2
private const val DATA_PLANE_WAIT_SLICE_MS = 50L
/** Connect to a TCP endpoint through the named EasyTier instance. */
@JvmStatic
suspend fun tcpConnect(
instanceName: String,
dstIp: String,
dstPort: Int,
timeoutMs: Long
): DataPlaneTcpStream {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneTcpConnectStart(
instanceName,
dstIp,
dstPort,
timeoutMs
)
)
val result = awaitOp(op) {
EasyTierDataPlaneJNI.dataPlaneTcpConnectFinish(it) ?: throw lastDataPlaneException()
}
return DataPlaneTcpStream(result.handle, result.localAddress)
}
/** Bind a TCP data-plane listener on [localPort]. Port 0 asks EasyTier to allocate one. */
@JvmStatic
suspend fun tcpBind(
instanceName: String,
localPort: Int,
timeoutMs: Long
): DataPlaneTcpListener {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneTcpBindStart(
instanceName,
localPort,
timeoutMs
)
)
val result = awaitOp(op) {
EasyTierDataPlaneJNI.dataPlaneTcpBindFinish(it) ?: throw lastDataPlaneException()
}
return DataPlaneTcpListener(result.handle, result.localAddress)
}
/** Accept one TCP stream from [listener]. */
@JvmStatic
suspend fun tcpAccept(listener: DataPlaneTcpListener, timeoutMs: Long): DataPlaneTcpStream {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneTcpAcceptStart(listener.handle, timeoutMs)
)
val result = awaitOp(op) {
EasyTierDataPlaneJNI.dataPlaneTcpAcceptFinish(it) ?: throw lastDataPlaneException()
}
return DataPlaneTcpStream(result.handle, result.localAddress, result.peerAddress)
}
/** Read up to [maxLength] bytes from [stream]. */
@JvmStatic
suspend fun tcpRead(
stream: DataPlaneTcpStream,
maxLength: Int,
timeoutMs: Long
): ByteArray {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneTcpReadStart(
stream.handle,
maxLength,
timeoutMs
)
)
return awaitOp(op) {
EasyTierDataPlaneJNI.dataPlaneTcpReadFinish(it)?.data
?: throw lastDataPlaneException()
}
}
/** Write [data] to [stream]. */
@JvmStatic
suspend fun tcpWrite(stream: DataPlaneTcpStream, data: ByteArray, timeoutMs: Long): Int {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneTcpWriteStart(
stream.handle,
data,
timeoutMs
)
)
return awaitOp(op) { EasyTierDataPlaneJNI.dataPlaneTcpWriteFinish(it) }
}
/** Bind a UDP data-plane socket on [localPort]. Port 0 asks EasyTier to allocate one. */
@JvmStatic
suspend fun udpBind(
instanceName: String,
localPort: Int,
timeoutMs: Long
): DataPlaneUdpSocket {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneUdpBindStart(
instanceName,
localPort,
timeoutMs
)
)
val result = awaitOp(op) {
EasyTierDataPlaneJNI.dataPlaneUdpBindFinish(it) ?: throw lastDataPlaneException()
}
return DataPlaneUdpSocket(result.handle, result.localAddress)
}
/** Send one UDP datagram through [socket]. */
@JvmStatic
suspend fun udpSendTo(
socket: DataPlaneUdpSocket,
dstIp: String,
dstPort: Int,
data: ByteArray,
timeoutMs: Long
): Int {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneUdpSendToStart(
socket.handle,
dstIp,
dstPort,
data,
timeoutMs
)
)
return awaitOp(op) { EasyTierDataPlaneJNI.dataPlaneUdpSendToFinish(it) }
}
/** Receive one UDP datagram through [socket]. */
@JvmStatic
suspend fun udpRecvFrom(
socket: DataPlaneUdpSocket,
maxLength: Int,
timeoutMs: Long
): DataPlaneUdpRecvResult {
val op =
requireOp(
EasyTierDataPlaneJNI.dataPlaneUdpRecvFromStart(
socket.handle,
maxLength,
timeoutMs
)
)
return awaitOp(op) {
EasyTierDataPlaneJNI.dataPlaneUdpRecvFromFinish(it) ?: throw lastDataPlaneException()
}
}
private fun requireOp(op: Long): Long {
if (op == 0L) {
throw lastDataPlaneException()
}
return op
}
private suspend fun <T> awaitOp(op: Long, finish: (Long) -> T): T =
withContext(Dispatchers.IO) {
var consumed = false
try {
awaitReady(op)
val result = finish(op)
consumed = true
result
} catch (e: CancellationException) {
EasyTierDataPlaneJNI.dataPlaneAsyncOpCancel(op)
throw e
} finally {
if (!consumed) {
EasyTierDataPlaneJNI.dataPlaneAsyncOpFree(op)
}
}
}
private suspend fun awaitReady(op: Long) {
while (true) {
currentCoroutineContext().ensureActive()
when (EasyTierDataPlaneJNI.dataPlaneAsyncOpWait(op, DATA_PLANE_WAIT_SLICE_MS)) {
DATA_PLANE_OP_READY, DATA_PLANE_OP_FAILED -> return
DATA_PLANE_OP_PENDING -> Unit
DATA_PLANE_OP_INVALID -> throw RuntimeException("Data-plane async operation is invalid")
else -> throw RuntimeException("Unknown data-plane async operation status")
}
}
}
private fun lastDataPlaneException(): RuntimeException {
return RuntimeException(EasyTierJNI.getLastError() ?: "EasyTier data-plane call failed")
}
}
@@ -1,8 +1,11 @@
package com.easytier.jni
/** EasyTier JNI 接口类 提供 Android 应用调用 EasyTier 网络功能的接口 */
object EasyTierJNI {
fun interface ConfigServerEventCallback {
fun onEvent(eventJson: String)
}
/** EasyTier JNI 接口类 提供 Android 应用调用 EasyTier 核心网络功能的接口 */
object EasyTierJNI {
init {
// 加载本地库
System.loadLibrary("easytier_android_jni")
@@ -33,6 +36,35 @@ object EasyTierJNI {
*/
@JvmStatic external fun runNetworkInstance(config: String): Int
/**
* 启动配置服务器客户端
* @param url 配置服务器 URL
* @param hostname 主机名,传入 null 使用系统主机名
* @param machineId 稳定机器 ID,由调用方负责持久化
* @param secureMode 是否启用 secure mode
* @param callback 远程配置应用/删除事件回调
* @return 0 表示成功,-1 表示失败
* @throws RuntimeException 当客户端启动失败时抛出异常
*/
@JvmStatic
external fun startConfigServerClient(
url: String,
hostname: String?,
machineId: String,
secureMode: Boolean,
callback: ConfigServerEventCallback?
): Int
/**
* 停止配置服务器客户端
* @return 0 表示成功,-1 表示失败
* @throws RuntimeException 当客户端停止失败时抛出异常
*/
@JvmStatic external fun stopConfigServerClient(): Int
/** 查询配置服务器客户端是否已连接 */
@JvmStatic external fun isConfigServerClientConnected(): Boolean
/**
* 保留指定的网络实例,停止其他实例
* @param instanceNames 要保留的实例名称数组,传入 null 或空数组将停止所有实例
@@ -44,11 +76,48 @@ object EasyTierJNI {
/**
* 收集网络信息
* @param maxLength 最大返回条目数
* @return 包含网络信息的字符串数组,每个元素格式为 "key=value"
* @return 包含网络信息的 JSON 字符串
* @throws RuntimeException 当操作失败时抛出异常
*/
@JvmStatic external fun collectNetworkInfos(maxLength: Int): String?
/**
* 列出当前运行的实例名称和实例 ID。
* @param maxLength 最大返回条目数
* @return JSON 对象,key 为 instance namevalue 为 instance id
* @throws RuntimeException 当操作失败时抛出异常
*/
@JvmStatic external fun listInstances(maxLength: Int): String?
/**
* 调用暴露的 EasyTier RPC 方法,输入和输出均为 protobuf JSON 字符串。
*
* 不支持 api.manage.WebClientService;实例启动、保留、删除、信息收集请继续使用专用 JNI API。
* payloadJson 需要包含目标 RPC 所需的 instance selector。
*
* @param serviceName RPC 服务名,例如 api.instance.PeerManageRpcService
* @param methodName RPC 方法名,支持 snake_case 或 proto 方法名
* @param domainName 仅 TcpProxyRpcService 使用;传 null 或空字符串默认 tcp
* @param payloadJson protobuf JSON 请求体
* @return protobuf JSON 响应体
* @throws RuntimeException 当 RPC 调用失败时抛出异常
*/
@JvmStatic
external fun callJsonRpc(
serviceName: String,
methodName: String,
domainName: String?,
payloadJson: String
): String?
/**
* 调用不需要 domainName 的 EasyTier RPC 方法。
*/
@JvmStatic
fun callJsonRpc(serviceName: String, methodName: String, payloadJson: String): String? {
return callJsonRpc(serviceName, methodName, null, payloadJson)
}
/**
* 获取最后的错误消息
* @return 错误消息字符串,如果没有错误则返回 null
@@ -0,0 +1,124 @@
use std::{
ffi::{CStr, c_char, c_void},
sync::{Arc, Mutex, MutexGuard},
};
use easytier_ffi::ConfigServerEventCallback;
use jni::JNIEnv;
use jni::objects::{GlobalRef, JObject, JValue};
use once_cell::sync::Lazy;
use crate::error;
pub(crate) struct JniConfigServerCallback {
java_vm: jni::JavaVM,
callback: GlobalRef,
}
static CONFIG_SERVER_CALLBACK: Lazy<Mutex<Option<Arc<JniConfigServerCallback>>>> =
Lazy::new(|| Mutex::new(None));
pub(crate) fn lock_callback_storage()
-> Result<MutexGuard<'static, Option<Arc<JniConfigServerCallback>>>, String> {
CONFIG_SERVER_CALLBACK
.lock()
.map_err(|e| format!("Failed to lock config server callback: {}", e))
}
pub(crate) fn new_callback(
env: &mut JNIEnv,
callback: &JObject,
) -> Result<Arc<JniConfigServerCallback>, String> {
let java_vm = env
.get_java_vm()
.map_err(|e| format!("Failed to get JavaVM: {:?}", e))?;
let callback = env
.new_global_ref(callback)
.map_err(|e| format!("Failed to create callback global ref: {:?}", e))?;
Ok(Arc::new(JniConfigServerCallback { java_vm, callback }))
}
pub(crate) fn callback_fn(
callback: &Option<Arc<JniConfigServerCallback>>,
) -> ConfigServerEventCallback {
callback
.as_ref()
.map(|_| config_server_event_callback as unsafe extern "C" fn(*const c_char, *mut c_void))
}
pub(crate) fn user_data(callback: &Option<Arc<JniConfigServerCallback>>) -> *mut c_void {
callback
.as_ref()
.map(|callback| Arc::as_ptr(callback) as *mut c_void)
.unwrap_or(std::ptr::null_mut())
}
impl JniConfigServerCallback {
fn clear_pending_exception(
env: &mut JNIEnv,
context: &str,
error: &dyn std::fmt::Debug,
) -> String {
match env.exception_check() {
Ok(true) => {
if let Err(clear_err) = env.exception_clear() {
return format!(
"{}: {:?}; failed to clear pending Java exception: {:?}",
context, error, clear_err
);
}
}
Ok(false) => {}
Err(check_err) => {
return format!(
"{}: {:?}; failed to check pending Java exception: {:?}",
context, error, check_err
);
}
}
format!("{}: {:?}", context, error)
}
fn on_event(&self, event_json: *const c_char) -> Result<(), String> {
let event_json = unsafe { CStr::from_ptr(event_json) }
.to_str()
.map_err(|e| format!("Invalid config server event JSON: {:?}", e))?;
let mut env = self
.java_vm
.attach_current_thread()
.map_err(|e| format!("Failed to attach callback thread: {:?}", e))?;
let event_json = env.new_string(event_json).map_err(|e| {
Self::clear_pending_exception(&mut env, "Failed to create event string", &e)
})?;
if let Err(e) = env.call_method(
self.callback.as_obj(),
"onEvent",
"(Ljava/lang/String;)V",
&[JValue::from(&event_json)],
) {
return Err(Self::clear_pending_exception(
&mut env,
"Failed to call config server callback",
&e,
));
}
Ok(())
}
}
unsafe extern "C" fn config_server_event_callback(
event_json: *const c_char,
user_data: *mut c_void,
) {
if event_json.is_null() || user_data.is_null() {
return;
}
let callback = unsafe { &*(user_data as *const JniConfigServerCallback) };
if let Err(error) = callback.on_event(event_json) {
error::set_callback_error(error);
}
}
@@ -0,0 +1,140 @@
use std::ptr;
use easytier_ffi::{
in_config_server_callback, is_config_server_client_connected, start_config_server_client,
stop_config_server_client,
};
use jni::JNIEnv;
use jni::objects::{JClass, JObject, JString};
use jni::sys::{JNI_FALSE, JNI_TRUE, jboolean, jint};
use crate::{
callback, error,
strings::{jstring_to_cstring, optional_jstring_to_cstring},
};
pub(crate) fn start_config_server_client_jni(
env: &mut JNIEnv,
config_server_url: JString,
hostname: JString,
machine_id: JString,
secure_mode: jboolean,
callback_obj: JObject,
) -> jint {
if in_config_server_callback() {
error::throw_exception(
env,
"Cannot start config server client from config server callback",
);
return -1;
}
let config_server_url = match jstring_to_cstring(env, &config_server_url) {
Ok(cstr) => cstr,
Err(e) => {
error::throw_exception(env, &format!("Invalid config server URL: {}", e));
return -1;
}
};
let hostname = match optional_jstring_to_cstring(env, &hostname) {
Ok(cstr) => cstr,
Err(e) => {
error::throw_exception(env, &format!("Invalid hostname: {}", e));
return -1;
}
};
let machine_id = match jstring_to_cstring(env, &machine_id) {
Ok(cstr) => cstr,
Err(e) => {
error::throw_exception(env, &format!("Invalid machine ID: {}", e));
return -1;
}
};
let callback_ref = if callback_obj.is_null() {
None
} else {
match callback::new_callback(env, &callback_obj) {
Ok(state) => Some(state),
Err(e) => {
error::throw_exception(env, &e);
return -1;
}
}
};
let mut callback_guard = match callback::lock_callback_storage() {
Ok(guard) => guard,
Err(e) => {
error::throw_exception(env, &e);
return -1;
}
};
if callback_guard.is_none() {
error::clear_callback_error();
}
let callback_fn = callback::callback_fn(&callback_ref);
let user_data = callback::user_data(&callback_ref);
let result = unsafe {
start_config_server_client(
config_server_url.as_ptr(),
hostname
.as_ref()
.map(|value| value.as_ptr())
.unwrap_or(ptr::null()),
machine_id.as_ptr(),
secure_mode == JNI_TRUE,
callback_fn,
user_data,
)
};
if result != 0 {
if let Some(error_msg) = error::get_last_error() {
error::throw_exception(env, &error_msg);
}
return result;
}
*callback_guard = callback_ref;
result
}
pub(crate) fn stop_config_server_client_jni(mut env: JNIEnv, _class: JClass) -> jint {
if in_config_server_callback() {
let result = stop_config_server_client();
if result != 0
&& let Some(error_msg) = error::get_last_error()
{
error::throw_exception(&mut env, &error_msg);
}
return result;
}
let mut callback_guard = match callback::lock_callback_storage() {
Ok(guard) => guard,
Err(e) => {
error::throw_exception(&mut env, &e);
return -1;
}
};
let result = stop_config_server_client();
if result != 0 {
if let Some(error_msg) = error::get_last_error() {
error::throw_exception(&mut env, &error_msg);
}
return result;
}
*callback_guard = None;
result
}
pub(crate) fn is_config_server_client_connected_jni(_env: JNIEnv, _class: JClass) -> jboolean {
if is_config_server_client_connected() != 0 {
JNI_TRUE
} else {
JNI_FALSE
}
}
@@ -0,0 +1,673 @@
use std::{
ffi::{CStr, c_char},
ptr,
};
use easytier_ffi::{
data_plane_async_op_cancel, data_plane_async_op_free, data_plane_async_op_status,
data_plane_async_op_wait, data_plane_free_bytes, data_plane_tcp_accept_finish,
data_plane_tcp_accept_start, data_plane_tcp_bind_finish, data_plane_tcp_bind_start,
data_plane_tcp_close, data_plane_tcp_connect_finish, data_plane_tcp_connect_start,
data_plane_tcp_listener_close, data_plane_tcp_read_finish, data_plane_tcp_read_start,
data_plane_tcp_write_finish, data_plane_tcp_write_start, data_plane_udp_bind_finish,
data_plane_udp_bind_start, data_plane_udp_close, data_plane_udp_recv_from_finish,
data_plane_udp_recv_from_start, data_plane_udp_send_to_finish, data_plane_udp_send_to_start,
free_string,
};
use jni::{
JNIEnv,
objects::{JByteArray, JClass, JObject, JString, JValue},
sys::{jint, jlong, jobject},
};
use crate::{
error::{get_last_error, throw_exception},
strings::jstring_to_cstring,
};
const SOCKET_ADDR_CLASS: &str = "com/easytier/jni/DataPlaneSocketAddress";
const TCP_CONNECT_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpConnectResult";
const TCP_BIND_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpBindResult";
const TCP_ACCEPT_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpAcceptResult";
const TCP_READ_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneTcpReadResult";
const UDP_BIND_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneUdpBindResult";
const UDP_RECV_RESULT_CLASS: &str = "com/easytier/jni/DataPlaneUdpRecvResult";
fn timeout_from_jlong(timeout_ms: jlong) -> u64 {
timeout_ms.max(0) as u64
}
fn port_from_jint(env: &mut JNIEnv, value: jint, name: &str) -> Option<u16> {
match u16::try_from(value) {
Ok(port) => Some(port),
Err(_) => {
throw_exception(env, &format!("Invalid {}: {}", name, value));
None
}
}
}
fn len_from_jint(env: &mut JNIEnv, value: jint, name: &str) -> Option<u32> {
match u32::try_from(value) {
Ok(len) => Some(len),
Err(_) => {
throw_exception(env, &format!("Invalid {}: {}", name, value));
None
}
}
}
fn throw_last(env: &mut JNIEnv) {
let message = get_last_error().unwrap_or_else(|| "EasyTier data-plane call failed".to_string());
throw_exception(env, &message);
}
unsafe fn take_ffi_string(ptr: *const c_char) -> String {
if ptr.is_null() {
return String::new();
}
let value = unsafe { CStr::from_ptr(ptr) }
.to_string_lossy()
.into_owned();
free_string(ptr);
value
}
fn new_socket_addr<'local>(
env: &mut JNIEnv<'local>,
ip: String,
port: u16,
) -> Option<JObject<'local>> {
let class = match env.find_class(SOCKET_ADDR_CLASS) {
Ok(class) => class,
Err(err) => {
throw_exception(
env,
&format!("Failed to find socket address class: {:?}", err),
);
return None;
}
};
let ip = match env.new_string(ip) {
Ok(ip) => ip,
Err(err) => {
throw_exception(env, &format!("Failed to create IP string: {:?}", err));
return None;
}
};
match env.new_object(
class,
"(Ljava/lang/String;I)V",
&[JValue::Object(&ip), JValue::Int(port as jint)],
) {
Ok(addr) => Some(addr),
Err(err) => {
throw_exception(env, &format!("Failed to create socket address: {:?}", err));
None
}
}
}
fn new_handle_addr_result(
env: &mut JNIEnv,
class_name: &str,
handle: u64,
ip: String,
port: u16,
) -> jobject {
let Some(addr) = new_socket_addr(env, ip, port) else {
return ptr::null_mut();
};
let class = match env.find_class(class_name) {
Ok(class) => class,
Err(err) => {
throw_exception(env, &format!("Failed to find result class: {:?}", err));
return ptr::null_mut();
}
};
let sig = format!("(JL{};)V", SOCKET_ADDR_CLASS);
match env.new_object(
class,
sig.as_str(),
&[JValue::Long(handle as jlong), JValue::Object(&addr)],
) {
Ok(result) => result.into_raw(),
Err(err) => {
throw_exception(env, &format!("Failed to create result object: {:?}", err));
ptr::null_mut()
}
}
}
fn close_tcp_stream_on_null(result: jobject, handle: u64) -> jobject {
if result.is_null() {
let _ = data_plane_tcp_close(handle);
}
result
}
fn close_tcp_listener_on_null(result: jobject, handle: u64) -> jobject {
if result.is_null() {
let _ = data_plane_tcp_listener_close(handle);
}
result
}
fn close_udp_socket_on_null(result: jobject, handle: u64) -> jobject {
if result.is_null() {
let _ = data_plane_udp_close(handle);
}
result
}
fn read_owned_bytes(ptr: *const u8, len: u32) -> Vec<u8> {
if ptr.is_null() || len == 0 {
return Vec::new();
}
let bytes = unsafe { std::slice::from_raw_parts(ptr, len as usize) }.to_vec();
data_plane_free_bytes(ptr, len);
bytes
}
pub(crate) fn async_op_status_jni(_env: JNIEnv, _class: JClass, handle: jlong) -> jint {
data_plane_async_op_status(handle as u64)
}
pub(crate) fn async_op_wait_jni(
_env: JNIEnv,
_class: JClass,
handle: jlong,
timeout_ms: jlong,
) -> jint {
data_plane_async_op_wait(handle as u64, timeout_ms.max(0) as u64)
}
pub(crate) fn async_op_cancel_jni(_env: JNIEnv, _class: JClass, handle: jlong) -> jint {
data_plane_async_op_cancel(handle as u64)
}
pub(crate) fn async_op_free_jni(_env: JNIEnv, _class: JClass, handle: jlong) -> jint {
data_plane_async_op_free(handle as u64)
}
pub(crate) fn tcp_connect_start_jni(
mut env: JNIEnv,
_class: JClass,
inst_name: JString,
dst_ip: JString,
dst_port: jint,
timeout_ms: jlong,
) -> jlong {
let inst_name = match jstring_to_cstring(&mut env, &inst_name) {
Ok(value) => value,
Err(err) => {
throw_exception(&mut env, &format!("Invalid instance name: {}", err));
return 0;
}
};
let dst_ip = match jstring_to_cstring(&mut env, &dst_ip) {
Ok(value) => value,
Err(err) => {
throw_exception(&mut env, &format!("Invalid destination IP: {}", err));
return 0;
}
};
let Some(dst_port) = port_from_jint(&mut env, dst_port, "destination port") else {
return 0;
};
let op = unsafe {
data_plane_tcp_connect_start(
inst_name.as_ptr(),
dst_ip.as_ptr(),
dst_port,
timeout_ms.max(0) as u64,
)
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn tcp_connect_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject {
let mut ip: *const c_char = ptr::null();
let mut port = 0u16;
let handle = unsafe { data_plane_tcp_connect_finish(op as u64, &mut ip, &mut port) };
if handle == 0 {
throw_last(&mut env);
return ptr::null_mut();
}
close_tcp_stream_on_null(
new_handle_addr_result(
&mut env,
TCP_CONNECT_RESULT_CLASS,
handle,
unsafe { take_ffi_string(ip) },
port,
),
handle,
)
}
pub(crate) fn tcp_bind_start_jni(
mut env: JNIEnv,
_class: JClass,
inst_name: JString,
local_port: jint,
timeout_ms: jlong,
) -> jlong {
let inst_name = match jstring_to_cstring(&mut env, &inst_name) {
Ok(value) => value,
Err(err) => {
throw_exception(&mut env, &format!("Invalid instance name: {}", err));
return 0;
}
};
let Some(local_port) = port_from_jint(&mut env, local_port, "local port") else {
return 0;
};
let op = unsafe {
data_plane_tcp_bind_start(
inst_name.as_ptr(),
local_port,
timeout_from_jlong(timeout_ms),
)
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn tcp_bind_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject {
let mut ip: *const c_char = ptr::null();
let mut port = 0u16;
let handle = unsafe { data_plane_tcp_bind_finish(op as u64, &mut ip, &mut port) };
if handle == 0 {
throw_last(&mut env);
return ptr::null_mut();
}
close_tcp_listener_on_null(
new_handle_addr_result(
&mut env,
TCP_BIND_RESULT_CLASS,
handle,
unsafe { take_ffi_string(ip) },
port,
),
handle,
)
}
pub(crate) fn tcp_accept_start_jni(
mut env: JNIEnv,
_class: JClass,
handle: jlong,
timeout_ms: jlong,
) -> jlong {
let op = unsafe { data_plane_tcp_accept_start(handle as u64, timeout_from_jlong(timeout_ms)) };
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn tcp_accept_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject {
let mut local_ip: *const c_char = ptr::null();
let mut local_port = 0u16;
let mut peer_ip: *const c_char = ptr::null();
let mut peer_port = 0u16;
let handle = unsafe {
data_plane_tcp_accept_finish(
op as u64,
&mut local_ip,
&mut local_port,
&mut peer_ip,
&mut peer_port,
)
};
if handle == 0 {
throw_last(&mut env);
return ptr::null_mut();
}
let Some(local_addr) =
new_socket_addr(&mut env, unsafe { take_ffi_string(local_ip) }, local_port)
else {
free_string(peer_ip);
let _ = data_plane_tcp_close(handle);
return ptr::null_mut();
};
let Some(peer_addr) = new_socket_addr(&mut env, unsafe { take_ffi_string(peer_ip) }, peer_port)
else {
let _ = data_plane_tcp_close(handle);
return ptr::null_mut();
};
let class = match env.find_class(TCP_ACCEPT_RESULT_CLASS) {
Ok(class) => class,
Err(err) => {
throw_exception(
&mut env,
&format!("Failed to find accept result class: {:?}", err),
);
let _ = data_plane_tcp_close(handle);
return ptr::null_mut();
}
};
let sig = format!("(JL{};L{};)V", SOCKET_ADDR_CLASS, SOCKET_ADDR_CLASS);
let result = match env.new_object(
class,
sig.as_str(),
&[
JValue::Long(handle as jlong),
JValue::Object(&local_addr),
JValue::Object(&peer_addr),
],
) {
Ok(result) => result.into_raw(),
Err(err) => {
throw_exception(
&mut env,
&format!("Failed to create accept result: {:?}", err),
);
ptr::null_mut()
}
};
close_tcp_stream_on_null(result, handle)
}
pub(crate) fn tcp_read_start_jni(
mut env: JNIEnv,
_class: JClass,
handle: jlong,
max_len: jint,
timeout_ms: jlong,
) -> jlong {
let Some(max_len) = len_from_jint(&mut env, max_len, "max length") else {
return 0;
};
let op = unsafe {
data_plane_tcp_read_start(handle as u64, max_len, timeout_from_jlong(timeout_ms))
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn tcp_read_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject {
let mut ptr: *const u8 = ptr::null();
let mut len = 0u32;
let ret = unsafe { data_plane_tcp_read_finish(op as u64, &mut ptr, &mut len) };
if ret < 0 {
throw_last(&mut env);
return ptr::null_mut();
}
let bytes = read_owned_bytes(ptr, len);
let array = match env.byte_array_from_slice(&bytes) {
Ok(array) => array,
Err(err) => {
throw_exception(&mut env, &format!("Failed to create byte array: {:?}", err));
return ptr::null_mut();
}
};
let class = match env.find_class(TCP_READ_RESULT_CLASS) {
Ok(class) => class,
Err(err) => {
throw_exception(
&mut env,
&format!("Failed to find read result class: {:?}", err),
);
return ptr::null_mut();
}
};
match env.new_object(class, "([B)V", &[JValue::Object(&array)]) {
Ok(result) => result.into_raw(),
Err(err) => {
throw_exception(
&mut env,
&format!("Failed to create read result: {:?}", err),
);
ptr::null_mut()
}
}
}
pub(crate) fn tcp_write_start_jni(
mut env: JNIEnv,
_class: JClass,
handle: jlong,
data: JByteArray,
timeout_ms: jlong,
) -> jlong {
let data = match env.convert_byte_array(&data) {
Ok(data) => data,
Err(err) => {
throw_exception(&mut env, &format!("Invalid write buffer: {:?}", err));
return 0;
}
};
let ptr = if data.is_empty() {
ptr::null()
} else {
data.as_ptr()
};
let op = unsafe {
data_plane_tcp_write_start(
handle as u64,
ptr,
data.len() as u32,
timeout_from_jlong(timeout_ms),
)
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn tcp_write_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jint {
let ret = data_plane_tcp_write_finish(op as u64);
if ret < 0 {
throw_last(&mut env);
}
ret
}
pub(crate) fn udp_bind_start_jni(
mut env: JNIEnv,
_class: JClass,
inst_name: JString,
local_port: jint,
timeout_ms: jlong,
) -> jlong {
let inst_name = match jstring_to_cstring(&mut env, &inst_name) {
Ok(value) => value,
Err(err) => {
throw_exception(&mut env, &format!("Invalid instance name: {}", err));
return 0;
}
};
let Some(local_port) = port_from_jint(&mut env, local_port, "local port") else {
return 0;
};
let op = unsafe {
data_plane_udp_bind_start(
inst_name.as_ptr(),
local_port,
timeout_from_jlong(timeout_ms),
)
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn udp_bind_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject {
let mut ip: *const c_char = ptr::null();
let mut port = 0u16;
let handle = unsafe { data_plane_udp_bind_finish(op as u64, &mut ip, &mut port) };
if handle == 0 {
throw_last(&mut env);
return ptr::null_mut();
}
close_udp_socket_on_null(
new_handle_addr_result(
&mut env,
UDP_BIND_RESULT_CLASS,
handle,
unsafe { take_ffi_string(ip) },
port,
),
handle,
)
}
pub(crate) fn udp_send_to_start_jni(
mut env: JNIEnv,
_class: JClass,
handle: jlong,
dst_ip: JString,
dst_port: jint,
data: JByteArray,
timeout_ms: jlong,
) -> jlong {
let dst_ip = match jstring_to_cstring(&mut env, &dst_ip) {
Ok(value) => value,
Err(err) => {
throw_exception(&mut env, &format!("Invalid destination IP: {}", err));
return 0;
}
};
let Some(dst_port) = port_from_jint(&mut env, dst_port, "destination port") else {
return 0;
};
let data = match env.convert_byte_array(&data) {
Ok(data) => data,
Err(err) => {
throw_exception(&mut env, &format!("Invalid UDP send buffer: {:?}", err));
return 0;
}
};
let ptr = if data.is_empty() {
ptr::null()
} else {
data.as_ptr()
};
let op = unsafe {
data_plane_udp_send_to_start(
handle as u64,
dst_ip.as_ptr(),
dst_port,
ptr,
data.len() as u32,
timeout_from_jlong(timeout_ms),
)
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn udp_send_to_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jint {
let ret = data_plane_udp_send_to_finish(op as u64);
if ret < 0 {
throw_last(&mut env);
}
ret
}
pub(crate) fn udp_recv_from_start_jni(
mut env: JNIEnv,
_class: JClass,
handle: jlong,
max_len: jint,
timeout_ms: jlong,
) -> jlong {
let Some(max_len) = len_from_jint(&mut env, max_len, "max length") else {
return 0;
};
let op = unsafe {
data_plane_udp_recv_from_start(handle as u64, max_len, timeout_from_jlong(timeout_ms))
};
if op == 0 {
throw_last(&mut env);
}
op as jlong
}
pub(crate) fn udp_recv_from_finish_jni(mut env: JNIEnv, _class: JClass, op: jlong) -> jobject {
let mut ptr: *const u8 = ptr::null();
let mut len = 0u32;
let mut ip: *const c_char = ptr::null();
let mut port = 0u16;
let ret = unsafe {
data_plane_udp_recv_from_finish(op as u64, &mut ptr, &mut len, &mut ip, &mut port)
};
if ret < 0 {
throw_last(&mut env);
return ptr::null_mut();
}
let bytes = read_owned_bytes(ptr, len);
let array = match env.byte_array_from_slice(&bytes) {
Ok(array) => array,
Err(err) => {
free_string(ip);
throw_exception(&mut env, &format!("Failed to create byte array: {:?}", err));
return ptr::null_mut();
}
};
let Some(peer_addr) = new_socket_addr(&mut env, unsafe { take_ffi_string(ip) }, port) else {
return ptr::null_mut();
};
let class = match env.find_class(UDP_RECV_RESULT_CLASS) {
Ok(class) => class,
Err(err) => {
throw_exception(
&mut env,
&format!("Failed to find UDP recv result class: {:?}", err),
);
return ptr::null_mut();
}
};
let sig = format!("([BL{};)V", SOCKET_ADDR_CLASS);
match env.new_object(
class,
sig.as_str(),
&[JValue::Object(&array), JValue::Object(&peer_addr)],
) {
Ok(result) => result.into_raw(),
Err(err) => {
throw_exception(
&mut env,
&format!("Failed to create UDP recv result: {:?}", err),
);
ptr::null_mut()
}
}
}
pub(crate) fn tcp_close_jni(mut env: JNIEnv, _class: JClass, handle: jlong) -> jint {
let ret = data_plane_tcp_close(handle as u64);
if ret != 0 {
throw_last(&mut env);
}
ret
}
pub(crate) fn tcp_listener_close_jni(mut env: JNIEnv, _class: JClass, handle: jlong) -> jint {
let ret = data_plane_tcp_listener_close(handle as u64);
if ret != 0 {
throw_last(&mut env);
}
ret
}
pub(crate) fn udp_close_jni(mut env: JNIEnv, _class: JClass, handle: jlong) -> jint {
let ret = data_plane_udp_close(handle as u64);
if ret != 0 {
throw_last(&mut env);
}
ret
}
@@ -0,0 +1,74 @@
use std::{
ffi::{CStr, c_char},
ptr,
sync::Mutex,
};
use easytier_ffi::{free_string, get_error_msg};
use jni::JNIEnv;
use jni::objects::JClass;
use jni::sys::jstring;
use once_cell::sync::Lazy;
static JNI_CALLBACK_ERROR: Lazy<Mutex<Option<String>>> = Lazy::new(|| Mutex::new(None));
pub(crate) fn set_callback_error(error: String) {
log::error!("{}", error);
if let Ok(mut guard) = JNI_CALLBACK_ERROR.lock() {
*guard = Some(error);
}
}
pub(crate) fn clear_callback_error() {
if let Ok(mut guard) = JNI_CALLBACK_ERROR.lock() {
*guard = None;
}
}
fn take_callback_error() -> Option<String> {
JNI_CALLBACK_ERROR
.lock()
.ok()
.and_then(|mut guard| guard.take())
}
fn get_ffi_last_error() -> Option<String> {
unsafe {
let mut error_ptr: *const c_char = ptr::null();
get_error_msg(&mut error_ptr);
if error_ptr.is_null() {
None
} else {
let error_cstr = CStr::from_ptr(error_ptr);
let error_str = error_cstr.to_string_lossy().into_owned();
free_string(error_ptr);
Some(error_str)
}
}
}
pub(crate) fn get_last_error() -> Option<String> {
match (get_ffi_last_error(), take_callback_error()) {
(Some(ffi_error), Some(callback_error)) => Some(format!(
"{}; config server callback error: {}",
ffi_error, callback_error
)),
(Some(ffi_error), None) => Some(ffi_error),
(None, Some(callback_error)) => Some(callback_error),
(None, None) => None,
}
}
pub(crate) fn throw_exception(env: &mut JNIEnv, message: &str) {
let _ = env.throw_new("java/lang/RuntimeException", message);
}
pub(crate) fn get_last_error_jni(env: JNIEnv, _class: JClass) -> jstring {
match get_last_error() {
Some(error) => match env.new_string(&error) {
Ok(jstr) => jstr.into_raw(),
Err(_) => ptr::null_mut(),
},
None => ptr::null_mut(),
}
}
@@ -0,0 +1,91 @@
use std::{
ffi::{CStr, c_char},
ptr,
};
use easytier_ffi::{call_json_rpc, free_string};
use jni::JNIEnv;
use jni::objects::{JClass, JString};
use jni::sys::jstring;
use crate::{
error::{get_last_error, throw_exception},
strings::{jstring_to_cstring, optional_jstring_to_cstring},
};
pub(crate) fn call_json_rpc_jni(
mut env: JNIEnv,
_class: JClass,
service_name: JString,
method_name: JString,
domain_name: JString,
payload_json: JString,
) -> jstring {
let service_name_cstr = match jstring_to_cstring(&mut env, &service_name) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid service name: {}", e));
return ptr::null_mut();
}
};
let method_name_cstr = match jstring_to_cstring(&mut env, &method_name) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid method name: {}", e));
return ptr::null_mut();
}
};
let domain_name_cstr = match optional_jstring_to_cstring(&mut env, &domain_name) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid domain name: {}", e));
return ptr::null_mut();
}
};
let payload_json_cstr = match jstring_to_cstring(&mut env, &payload_json) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid payload JSON: {}", e));
return ptr::null_mut();
}
};
let domain_name_ptr = domain_name_cstr
.as_ref()
.map_or(ptr::null(), |cstr| cstr.as_ptr());
let mut response_ptr: *const c_char = ptr::null();
let result = unsafe {
call_json_rpc(
service_name_cstr.as_ptr(),
method_name_cstr.as_ptr(),
domain_name_ptr,
payload_json_cstr.as_ptr(),
&mut response_ptr,
)
};
if result != 0 {
if let Some(error) = get_last_error() {
throw_exception(&mut env, &error);
}
return ptr::null_mut();
}
if response_ptr.is_null() {
throw_exception(&mut env, "JSON RPC returned a null response");
return ptr::null_mut();
}
let response = unsafe { CStr::from_ptr(response_ptr) }
.to_string_lossy()
.into_owned();
free_string(response_ptr);
match env.new_string(&response) {
Ok(jstr) => jstr.into_raw(),
Err(_) => {
throw_exception(&mut env, "Failed to create JSON RPC response string");
ptr::null_mut()
}
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,13 @@
use once_cell::sync::Lazy;
static LOGGER_INIT: Lazy<()> = Lazy::new(|| {
android_logger::init_once(
android_logger::Config::default()
.with_max_level(log::LevelFilter::Debug)
.with_tag("EasyTier-JNI"),
);
});
pub(crate) fn init() {
Lazy::force(&LOGGER_INIT);
}
@@ -0,0 +1,261 @@
use std::{ffi::CStr, ptr};
use easytier::proto::api::manage::{NetworkInstanceRunningInfo, NetworkInstanceRunningInfoMap};
use easytier_ffi::{
KeyValuePair, collect_network_infos, free_string, list_instance, parse_config,
retain_network_instance, run_network_instance, set_tun_fd,
};
use jni::JNIEnv;
use jni::objects::{JClass, JObjectArray, JString};
use jni::sys::{jint, jstring};
use crate::{
error::{get_last_error, throw_exception},
strings::jstring_to_cstring,
};
pub(crate) fn set_tun_fd_jni(
mut env: JNIEnv,
_class: JClass,
inst_name: JString,
fd: jint,
) -> jint {
let inst_name_cstr = match jstring_to_cstring(&mut env, &inst_name) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid instance name: {}", e));
return -1;
}
};
unsafe {
let result = set_tun_fd(inst_name_cstr.as_ptr(), fd);
if result != 0
&& let Some(error) = get_last_error()
{
throw_exception(&mut env, &error);
}
result
}
}
pub(crate) fn parse_config_jni(mut env: JNIEnv, _class: JClass, config: JString) -> jint {
let config_cstr = match jstring_to_cstring(&mut env, &config) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid config string: {}", e));
return -1;
}
};
unsafe {
let result = parse_config(config_cstr.as_ptr());
if result != 0
&& let Some(error) = get_last_error()
{
throw_exception(&mut env, &error);
}
result
}
}
pub(crate) fn run_network_instance_jni(mut env: JNIEnv, _class: JClass, config: JString) -> jint {
let config_cstr = match jstring_to_cstring(&mut env, &config) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(&mut env, &format!("Invalid config string: {}", e));
return -1;
}
};
unsafe {
let result = run_network_instance(config_cstr.as_ptr());
if result != 0
&& let Some(error) = get_last_error()
{
throw_exception(&mut env, &error);
}
result
}
}
pub(crate) fn retain_network_instance_jni(
mut env: JNIEnv,
_class: JClass,
instance_names: JObjectArray,
) -> jint {
if instance_names.is_null() {
return retain_all(&mut env);
}
let array_length = match env.get_array_length(&instance_names) {
Ok(len) => len as usize,
Err(e) => {
throw_exception(&mut env, &format!("Failed to get array length: {:?}", e));
return -1;
}
};
if array_length == 0 {
return retain_all(&mut env);
}
let mut c_strings = Vec::with_capacity(array_length);
let mut c_string_ptrs = Vec::with_capacity(array_length);
for i in 0..array_length {
let java_string = match env.get_object_array_element(&instance_names, i as i32) {
Ok(obj) => obj,
Err(e) => {
throw_exception(
&mut env,
&format!("Failed to get array element {}: {:?}", i, e),
);
return -1;
}
};
if java_string.is_null() {
throw_exception(
&mut env,
&format!("Invalid instance name at index {}: null", i),
);
return -1;
}
let jstring = JString::from(java_string);
let c_string = match jstring_to_cstring(&mut env, &jstring) {
Ok(cstr) => cstr,
Err(e) => {
throw_exception(
&mut env,
&format!("Invalid instance name at index {}: {}", i, e),
);
return -1;
}
};
c_string_ptrs.push(c_string.as_ptr());
c_strings.push(c_string);
}
unsafe {
let result = retain_network_instance(c_string_ptrs.as_ptr(), c_string_ptrs.len());
if result != 0
&& let Some(error) = get_last_error()
{
throw_exception(&mut env, &error);
}
result
}
}
fn retain_all(env: &mut JNIEnv) -> jint {
unsafe {
let result = retain_network_instance(ptr::null(), 0);
if result != 0
&& let Some(error) = get_last_error()
{
throw_exception(env, &error);
}
result
}
}
pub(crate) fn collect_network_infos_jni(
mut env: JNIEnv,
_class: JClass,
max_length: jint,
) -> jstring {
let max_length = max_length.max(0) as usize;
let mut infos = vec![
KeyValuePair {
key: ptr::null(),
value: ptr::null(),
};
max_length
];
unsafe {
let count = collect_network_infos(infos.as_mut_ptr(), max_length);
if count < 0 {
if let Some(error) = get_last_error() {
throw_exception(&mut env, &error);
}
return ptr::null_mut();
}
let mut ret = NetworkInstanceRunningInfoMap::default();
for info in infos.iter().take(count as usize) {
let key_ptr = info.key;
let val_ptr = info.value;
if key_ptr.is_null() || val_ptr.is_null() {
break;
}
let key = CStr::from_ptr(key_ptr).to_string_lossy().into_owned();
let val = CStr::from_ptr(val_ptr).to_string_lossy().into_owned();
free_string(key_ptr);
free_string(val_ptr);
let value = match serde_json::from_str::<NetworkInstanceRunningInfo>(&val) {
Ok(v) => v,
Err(_) => {
throw_exception(&mut env, "Failed to parse JSON");
continue;
}
};
ret.map.insert(key, value);
}
let json_str = serde_json::to_string(&ret).unwrap_or_else(|_| "{}".to_string());
match env.new_string(&json_str) {
Ok(jstr) => jstr.into_raw(),
Err(_) => {
throw_exception(&mut env, "Failed to create JSON string");
ptr::null_mut()
}
}
}
}
pub(crate) fn list_instances_jni(mut env: JNIEnv, _class: JClass, max_length: jint) -> jstring {
let max_length = max_length.max(0) as usize;
let mut infos = vec![
KeyValuePair {
key: ptr::null(),
value: ptr::null(),
};
max_length
];
unsafe {
let count = list_instance(infos.as_mut_ptr(), max_length);
if count < 0 {
if let Some(error) = get_last_error() {
throw_exception(&mut env, &error);
}
return ptr::null_mut();
}
let mut ret = serde_json::Map::new();
for info in infos.iter().take(count as usize) {
let key_ptr = info.key;
let val_ptr = info.value;
if key_ptr.is_null() || val_ptr.is_null() {
break;
}
let key = CStr::from_ptr(key_ptr).to_string_lossy().into_owned();
let val = CStr::from_ptr(val_ptr).to_string_lossy().into_owned();
free_string(key_ptr);
free_string(val_ptr);
ret.insert(key, serde_json::Value::String(val));
}
let json_str = serde_json::Value::Object(ret).to_string();
match env.new_string(&json_str) {
Ok(jstr) => jstr.into_raw(),
Err(_) => {
throw_exception(&mut env, "Failed to create instance list JSON string");
ptr::null_mut()
}
}
}
}
@@ -0,0 +1,23 @@
use std::ffi::CString;
use jni::JNIEnv;
use jni::objects::JString;
pub(crate) fn jstring_to_cstring(env: &mut JNIEnv, jstr: &JString) -> Result<CString, String> {
let java_str = env
.get_string(jstr)
.map_err(|e| format!("Failed to get string: {:?}", e))?;
let rust_str = java_str.to_str().map_err(|_| "Invalid UTF-8".to_string())?;
CString::new(rust_str).map_err(|_| "String contains null byte".to_string())
}
pub(crate) fn optional_jstring_to_cstring(
env: &mut JNIEnv,
jstr: &JString,
) -> Result<Option<CString>, String> {
if jstr.is_null() {
return Ok(None);
}
jstring_to_cstring(env, jstr).map(Some)
}
+12 -1
View File
@@ -4,14 +4,25 @@ version = "0.1.0"
edition.workspace = true
[lib]
crate-type = ["cdylib"]
crate-type = ["cdylib", "rlib"]
[features]
default = ["c-abi", "ffi-dataplane"]
c-abi = []
ffi-dataplane = ["easytier/ffi-dataplane"]
[dependencies]
easytier = { path = "../../easytier" }
once_cell = "1.18.0"
dashmap = "6.0"
tokio = { version = "1", features = ["rt-multi-thread", "io-util", "time", "sync", "macros"] }
async-trait = "0.1"
log = "0.4"
percent-encoding = "2.3"
url = "2"
serde = { version = "1.0", features = ["derive"] }
serde_json = "1"
uuid = "1.17.0"
tokio-util = "0.7"
@@ -0,0 +1,429 @@
#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#define DATA_PLANE_OP_PENDING 0
#define DATA_PLANE_OP_READY 1
#define DATA_PLANE_OP_FAILED -1
#define DATA_PLANE_OP_INVALID -2
extern int run_network_instance(const char *cfg_str);
extern void get_error_msg(const char **out);
extern void free_string(const char *s);
extern int data_plane_async_op_status(uint64_t op);
extern int data_plane_async_op_wait(uint64_t op, uint64_t timeout_ms);
extern int data_plane_async_op_cancel(uint64_t op);
extern int data_plane_async_op_free(uint64_t op);
extern void data_plane_free_bytes(const uint8_t *ptr, uint32_t len);
extern uint64_t data_plane_tcp_connect_start(
const char *inst_name,
const char *dst_ip,
uint16_t dst_port,
uint64_t timeout_ms);
extern uint64_t data_plane_tcp_connect_finish(
uint64_t op,
const char **out_local_ip,
uint16_t *out_local_port);
extern uint64_t data_plane_tcp_bind_start(
const char *inst_name,
uint16_t local_port,
uint64_t timeout_ms);
extern uint64_t data_plane_tcp_bind_finish(
uint64_t op,
const char **out_local_ip,
uint16_t *out_local_port);
extern uint64_t data_plane_tcp_accept_start(uint64_t listener, uint64_t timeout_ms);
extern uint64_t data_plane_tcp_accept_finish(
uint64_t op,
const char **out_local_ip,
uint16_t *out_local_port,
const char **out_peer_ip,
uint16_t *out_peer_port);
extern uint64_t data_plane_tcp_read_start(
uint64_t stream,
uint32_t max_len,
uint64_t timeout_ms);
extern int data_plane_tcp_read_finish(
uint64_t op,
const uint8_t **out_buf,
uint32_t *out_len);
extern uint64_t data_plane_tcp_write_start(
uint64_t stream,
const uint8_t *buf,
uint32_t len,
uint64_t timeout_ms);
extern int data_plane_tcp_write_finish(uint64_t op);
extern int data_plane_tcp_close(uint64_t stream);
extern int data_plane_tcp_listener_close(uint64_t listener);
extern uint64_t data_plane_udp_bind_start(
const char *inst_name,
uint16_t local_port,
uint64_t timeout_ms);
extern uint64_t data_plane_udp_bind_finish(
uint64_t op,
const char **out_local_ip,
uint16_t *out_local_port);
extern uint64_t data_plane_udp_send_to_start(
uint64_t socket,
const char *dst_ip,
uint16_t dst_port,
const uint8_t *buf,
uint32_t len,
uint64_t timeout_ms);
extern int data_plane_udp_send_to_finish(uint64_t op);
extern uint64_t data_plane_udp_recv_from_start(
uint64_t socket,
uint32_t max_len,
uint64_t timeout_ms);
extern int data_plane_udp_recv_from_finish(
uint64_t op,
const uint8_t **out_buf,
uint32_t *out_len,
const char **out_ip,
uint16_t *out_port);
extern int data_plane_udp_close(uint64_t socket);
static void print_last_error(const char *prefix) {
const char *err = NULL;
get_error_msg(&err);
if (err) {
fprintf(stderr, "%s: %s\n", prefix, err);
free_string(err);
} else {
fprintf(stderr, "%s\n", prefix);
}
}
static int parse_ip_port(const char *value, char *ip, size_t ip_len, uint16_t *port) {
const char *colon = strrchr(value, ':');
if (!colon || colon == value || !colon[1]) {
fprintf(stderr, "expected IPv4 target in IP:PORT form, got %s\n", value);
return -1;
}
size_t host_len = (size_t)(colon - value);
if (host_len >= ip_len) {
fprintf(stderr, "IP address is too long: %s\n", value);
return -1;
}
char *end = NULL;
long parsed_port = strtol(colon + 1, &end, 10);
if (!end || *end != '\0' || parsed_port < 0 || parsed_port > 65535) {
fprintf(stderr, "invalid port in %s\n", value);
return -1;
}
memcpy(ip, value, host_len);
ip[host_len] = '\0';
*port = (uint16_t)parsed_port;
return 0;
}
static int wait_op(uint64_t op, uint64_t timeout_ms) {
uint64_t waited = 0;
while (waited < timeout_ms) {
int status = data_plane_async_op_wait(op, 50);
if (status != DATA_PLANE_OP_PENDING) {
return status;
}
waited += 50;
}
return data_plane_async_op_status(op);
}
static int wait_or_cancel(uint64_t op, uint64_t timeout_ms, const char *what) {
int status = wait_op(op, timeout_ms);
if (status == DATA_PLANE_OP_READY || status == DATA_PLANE_OP_FAILED) {
return status;
}
if (status == DATA_PLANE_OP_PENDING) {
fprintf(stderr, "%s did not finish within %llu ms\n", what, (unsigned long long)timeout_ms);
data_plane_async_op_cancel(op);
data_plane_async_op_free(op);
return DATA_PLANE_OP_INVALID;
}
fprintf(stderr, "%s returned invalid op status %d\n", what, status);
return status;
}
static int async_tcp_read_once(uint64_t stream, uint64_t timeout_ms) {
uint64_t op = data_plane_tcp_read_start(stream, 512, timeout_ms);
if (!op) {
print_last_error("tcp read start failed");
return -1;
}
if (wait_or_cancel(op, timeout_ms + 1000, "tcp read") == DATA_PLANE_OP_INVALID) {
return -1;
}
const uint8_t *buf = NULL;
uint32_t len = 0;
int ret = data_plane_tcp_read_finish(op, &buf, &len);
if (ret < 0) {
print_last_error("tcp read finish failed");
return -1;
}
printf("tcp read %d bytes: %.*s\n", ret, ret, buf ? (const char *)buf : "");
data_plane_free_bytes(buf, len);
return 0;
}
static int async_tcp_write_all(uint64_t stream, const char *data, uint64_t timeout_ms) {
uint64_t op = data_plane_tcp_write_start(
stream,
(const uint8_t *)data,
(uint32_t)strlen(data),
timeout_ms);
if (!op) {
print_last_error("tcp write start failed");
return -1;
}
if (wait_or_cancel(op, timeout_ms + 1000, "tcp write") == DATA_PLANE_OP_INVALID) {
return -1;
}
int ret = data_plane_tcp_write_finish(op);
if (ret < 0) {
print_last_error("tcp write finish failed");
return -1;
}
printf("tcp wrote %d bytes\n", ret);
return 0;
}
static int run_tcp_connect_demo(const char *inst, const char *target) {
char ip[128];
uint16_t port = 0;
if (parse_ip_port(target, ip, sizeof(ip), &port) != 0) {
return -1;
}
uint64_t op = data_plane_tcp_connect_start(inst, ip, port, 30000);
if (!op) {
print_last_error("tcp connect start failed");
return -1;
}
if (wait_or_cancel(op, 31000, "tcp connect") == DATA_PLANE_OP_INVALID) {
return -1;
}
const char *local_ip = NULL;
uint16_t local_port = 0;
uint64_t stream = data_plane_tcp_connect_finish(op, &local_ip, &local_port);
if (!stream) {
print_last_error("tcp connect finish failed");
return -1;
}
printf("tcp connected from %s:%u to %s:%u, handle=%llu\n",
local_ip,
local_port,
ip,
port,
(unsigned long long)stream);
free_string(local_ip);
int ret = async_tcp_read_once(stream, 10000);
data_plane_tcp_close(stream);
return ret;
}
static int run_tcp_listen_demo(const char *inst, const char *port_text) {
uint16_t port = (uint16_t)strtoul(port_text, NULL, 10);
uint64_t op = data_plane_tcp_bind_start(inst, port, 30000);
if (!op) {
print_last_error("tcp bind start failed");
return -1;
}
if (wait_or_cancel(op, 31000, "tcp bind") == DATA_PLANE_OP_INVALID) {
return -1;
}
const char *local_ip = NULL;
uint16_t local_port = 0;
uint64_t listener = data_plane_tcp_bind_finish(op, &local_ip, &local_port);
if (!listener) {
print_last_error("tcp bind finish failed");
return -1;
}
printf("tcp listening on %s:%u, handle=%llu\n",
local_ip,
local_port,
(unsigned long long)listener);
free_string(local_ip);
op = data_plane_tcp_accept_start(listener, 60000);
if (!op) {
print_last_error("tcp accept start failed");
data_plane_tcp_listener_close(listener);
return -1;
}
if (wait_or_cancel(op, 61000, "tcp accept") == DATA_PLANE_OP_INVALID) {
data_plane_tcp_listener_close(listener);
return -1;
}
const char *peer_ip = NULL;
uint16_t peer_port = 0;
local_ip = NULL;
local_port = 0;
uint64_t stream = data_plane_tcp_accept_finish(
op,
&local_ip,
&local_port,
&peer_ip,
&peer_port);
data_plane_tcp_listener_close(listener);
if (!stream) {
print_last_error("tcp accept finish failed");
return -1;
}
printf("tcp accepted %s:%u -> %s:%u, stream=%llu\n",
peer_ip,
peer_port,
local_ip,
local_port,
(unsigned long long)stream);
free_string(local_ip);
free_string(peer_ip);
int ret = async_tcp_read_once(stream, 10000);
if (ret == 0) {
ret = async_tcp_write_all(stream, "pong", 10000);
}
data_plane_tcp_close(stream);
return ret;
}
static int run_udp_demo(const char *inst, const char *target) {
char ip[128];
uint16_t port = 0;
if (parse_ip_port(target, ip, sizeof(ip), &port) != 0) {
return -1;
}
uint64_t op = data_plane_udp_bind_start(inst, 0, 30000);
if (!op) {
print_last_error("udp bind start failed");
return -1;
}
if (wait_or_cancel(op, 31000, "udp bind") == DATA_PLANE_OP_INVALID) {
return -1;
}
const char *local_ip = NULL;
uint16_t local_port = 0;
uint64_t socket = data_plane_udp_bind_finish(op, &local_ip, &local_port);
if (!socket) {
print_last_error("udp bind finish failed");
return -1;
}
printf("udp bound on %s:%u, handle=%llu\n",
local_ip,
local_port,
(unsigned long long)socket);
free_string(local_ip);
const char payload[] = "ping";
op = data_plane_udp_send_to_start(
socket,
ip,
port,
(const uint8_t *)payload,
(uint32_t)strlen(payload),
10000);
if (!op) {
print_last_error("udp send start failed");
data_plane_udp_close(socket);
return -1;
}
if (wait_or_cancel(op, 11000, "udp send") == DATA_PLANE_OP_INVALID) {
data_plane_udp_close(socket);
return -1;
}
int sent = data_plane_udp_send_to_finish(op);
if (sent < 0) {
print_last_error("udp send finish failed");
data_plane_udp_close(socket);
return -1;
}
printf("udp sent %d bytes to %s:%u\n", sent, ip, port);
op = data_plane_udp_recv_from_start(socket, 512, 30000);
if (!op) {
print_last_error("udp recv start failed");
data_plane_udp_close(socket);
return -1;
}
if (wait_or_cancel(op, 31000, "udp recv") == DATA_PLANE_OP_INVALID) {
data_plane_udp_close(socket);
return -1;
}
const uint8_t *buf = NULL;
uint32_t len = 0;
const char *peer_ip = NULL;
uint16_t peer_port = 0;
int ret = data_plane_udp_recv_from_finish(op, &buf, &len, &peer_ip, &peer_port);
if (ret < 0) {
print_last_error("udp recv finish failed");
data_plane_udp_close(socket);
return -1;
}
printf("udp received %d bytes from %s:%u: %.*s\n",
ret,
peer_ip,
peer_port,
ret,
buf ? (const char *)buf : "");
data_plane_free_bytes(buf, len);
free_string(peer_ip);
data_plane_udp_close(socket);
return 0;
}
static void print_usage(void) {
printf("Set EASYTIER_FFI_CONFIG and EASYTIER_FFI_INSTANCE to run the async data-plane demo.\n");
printf("Optional demos:\n");
printf(" EASYTIER_FFI_TARGET=10.0.0.2:22 async TCP connect/read\n");
printf(" EASYTIER_FFI_LISTEN_PORT=12345 async TCP bind/accept/read/write\n");
printf(" EASYTIER_FFI_UDP_TARGET=10.0.0.2:9000 async UDP bind/send_to/recv_from\n");
}
int main(void) {
const char *config = getenv("EASYTIER_FFI_CONFIG");
const char *instance = getenv("EASYTIER_FFI_INSTANCE");
if (!config || !instance) {
print_usage();
return 0;
}
if (run_network_instance(config) != 0) {
print_last_error("run_network_instance failed");
return 1;
}
printf("network instance started: %s\n", instance);
int failed = 0;
const char *target = getenv("EASYTIER_FFI_TARGET");
if (target) {
failed |= run_tcp_connect_demo(instance, target) != 0;
}
const char *listen_port = getenv("EASYTIER_FFI_LISTEN_PORT");
if (listen_port) {
failed |= run_tcp_listen_demo(instance, listen_port) != 0;
}
const char *udp_target = getenv("EASYTIER_FFI_UDP_TARGET");
if (udp_target) {
failed |= run_udp_demo(instance, udp_target) != 0;
}
if (!target && !listen_port && !udp_target) {
printf("No dataplane demo env var was set; nothing else to run.\n");
print_usage();
}
return failed ? 1 : 0;
}
@@ -0,0 +1,100 @@
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <stdbool.h>
#include <unistd.h> // for sleep
// FFI struct and function declarations
typedef struct {
const char* key;
const char* value;
} KeyValuePair;
typedef void (*config_server_event_callback)(
const char* event_json,
void* user_data
);
extern int parse_config(const char* cfg_str);
extern int run_network_instance(const char* cfg_str);
extern void get_error_msg(const char** out);
extern void free_string(const char* s);
extern int collect_network_infos(KeyValuePair* infos, size_t max_length);
extern int start_config_server_client(
const char* config_server_url,
const char* hostname,
const char* machine_id,
bool secure_mode,
config_server_event_callback callback,
void* user_data
);
extern int stop_config_server_client(void);
extern int is_config_server_client_connected(void);
static void on_config_server_event(const char* event_json, void* user_data) {
(void)user_data;
printf("config server event: %s\n", event_json);
}
int main() {
const char* config = "inst_name = \"test\"\nnetwork = \"test_network\"\n";
int ret;
// 调用 parse_config
ret = parse_config(config);
if (ret != 0) {
const char* err = NULL;
get_error_msg(&err);
if (err) {
printf("parse_config error: %s\n", err);
free_string(err);
}
return 1;
}
printf("parse_config success\n");
// 调用 run_network_instance
ret = run_network_instance(config);
if (ret != 0) {
const char* err = NULL;
get_error_msg(&err);
if (err) {
printf("run_network_instance error: %s\n", err);
free_string(err);
}
return 1;
}
printf("run_network_instance success\n");
// 周期性调用 collect_network_infos 并打印
const size_t max_infos = 8;
KeyValuePair* infos = (KeyValuePair*)malloc(sizeof(KeyValuePair) * max_infos);
if (!infos) {
fprintf(stderr, "malloc failed\n");
return 1;
}
for (int i = 0; i < 5; ++i) { // 循环5次作为示例
memset(infos, 0, sizeof(KeyValuePair) * max_infos);
int count = collect_network_infos(infos, max_infos);
if (count < 0) {
const char* err = NULL;
get_error_msg(&err);
if (err) {
printf("collect_network_infos error: %s\n", err);
free_string(err);
}
break;
}
printf("collect_network_infos: %d instance(s)\n", count);
for (int j = 0; j < count; ++j) {
printf(" [%d] key: %s\n value: %s\n", j, infos[j].key, infos[j].value);
free_string(infos[j].key);
free_string(infos[j].value);
}
sleep(1);
}
free(infos);
return 0;
}
@@ -0,0 +1,138 @@
# 1. Go FFI Demo
This demo wraps EasyTier FFI data-plane TCP as Go `net.Conn` and `net.Listener`.
It can connect to an SSH server through EasyTier and read its banner, or accept a
TCP connection from another EasyTier peer and run a small ping/pong exchange.
The async op-handle wrapper is in `easytier_async.go`; the original synchronous
wrapper stays in `easytier.go`.
## 1.1. Build the FFI library
Run from the repository root:
```sh
cargo build -p easytier-ffi --features ffi-dataplane
```
The demo loads the debug library by default:
```text
target/debug/libeasytier_ffi.so
```
To use another library path, export `EASYTIER_FFI_LIB=/path/to/libeasytier_ffi.so`.
## 1.2. Configure the EasyTier config
`EASYTIER_FFI_CONFIG` is a string of the EasyTier config in TOML format which is passed to the FFI library. For example:
```sh
export EASYTIER_FFI_CONFIG='instance_name = "default"
ipv4 = "10.0.0.1"
[network_identity]
network_name = "testnet"
network_secret = "mysecret"
[flags]
no_tun = true # disable tun device to avoid permission issues.
bind_device = false # allow loopback peers in local examples.
[[peer]]
uri = "tcp://123.123.123.123:11010"
'
```
You should configure with your own real values.
Set the local instance name and a SSH server target to connect through EasyTier:
```sh
export EASYTIER_FFI_INSTANCE=default
export EASYTIER_FFI_TARGET=10.0.0.2:22
```
To run the TCP listen integration test in the same `go test` process as the SSH
test, use a separate instance name and config:
```sh
export EASYTIER_FFI_LISTEN_CONFIG='instance_name = "listener"
ipv4 = "10.0.0.3"
[network_identity]
network_name = "testnet"
network_secret = "mysecret"
[flags]
no_tun = true
bind_device = false
[[peer]]
uri = "tcp://123.123.123.123:11010"
'
export EASYTIER_FFI_LISTEN_INSTANCE=listener
export EASYTIER_FFI_LISTEN_PORT=12345
```
## 1.3. Run the demo
`goffi` is built without cgo on Linux, so run the tests with `CGO_ENABLED=0`:
```sh
cd easytier-contrib/easytier-ffi/examples/go
CGO_ENABLED=0 go test -v ./...
```
The synchronous tests use the environment variables above. The async Go tests
are self-contained: they start two local EasyTier instances in the same test
process with `no_tun = true` and `bind_device = false`, then run TCP and UDP
ping/pong over the async data-plane API.
The synchronous wrapper also exposes `CallJSONRPC(service, method, domain,
payload)` for non-lifecycle EasyTier RPCs. For example,
`CallJSONRPC("api.logger.LoggerRpcService", "get_logger_config", "", "{}")`
returns the logger config as protobuf JSON. Instance lifecycle management RPCs
are intentionally filtered; use the dedicated FFI APIs for starting and
stopping instances.
To run only the async tests:
```sh
cd easytier-contrib/easytier-ffi/examples/go
CGO_ENABLED=0 go test -run 'TestAsync' -v ./...
```
When the SSH integration environment variables are set, expected synchronous
test output includes an SSH banner similar to:
```text
attempt 1: got banner "SSH-2.0-..."
PASS
```
For `TestTCPListenIntegration`, connect from another EasyTier peer to the local
EasyTier IPv4 address and `EASYTIER_FFI_LISTEN_PORT`, send `ping`, and expect
`pong` in response.
The async test output should include local TCP bind/connect log lines and finish
with `PASS` without any extra environment variables.
## 1.4. C async example
The C async example is kept separate from the basic C example:
```sh
cargo build -p easytier-ffi --features ffi-dataplane
cc -Wall -Wextra -pedantic \
../example_data_plane_async.c \
-L ../../../../target/debug -leasytier_ffi \
-Wl,-rpath,../../../../target/debug \
-o /tmp/easytier_data_plane_async
/tmp/easytier_data_plane_async
```
Without environment variables it prints usage and exits successfully. With
`EASYTIER_FFI_CONFIG`, `EASYTIER_FFI_INSTANCE`, and one of
`EASYTIER_FFI_TARGET`, `EASYTIER_FFI_LISTEN_PORT`, or `EASYTIER_FFI_UDP_TARGET`,
it runs the corresponding async data-plane flow.
@@ -0,0 +1,593 @@
package easytierffi
import (
"context"
"errors"
"fmt"
"io"
"net"
"os"
"runtime"
"strconv"
"strings"
"sync/atomic"
"time"
"unsafe"
"github.com/go-webgpu/goffi/ffi"
"github.com/go-webgpu/goffi/types"
)
const defaultTimeout = 30 * time.Second
type Native struct {
lib unsafe.Pointer
runNetworkInstance symCall
callJSONRPC symCall
getErrorMsg symCall
freeString symCall
tcpConnect symCall
tcpBind symCall
tcpAccept symCall
tcpRead symCall
tcpWrite symCall
tcpClose symCall
tcpListenerClose symCall
}
type Conn struct {
native *Native
handle uint64
local net.Addr
remote net.Addr
closed atomic.Bool
rd atomicDeadline
wd atomicDeadline
}
type Listener struct {
native *Native
handle uint64
addr net.Addr
closed atomic.Bool
}
type symCall struct {
fn unsafe.Pointer
cif types.CallInterface
}
type atomicDeadline struct{ v atomic.Int64 }
type timeoutError string
func Open(path string) (*Native, error) {
lib, err := ffi.LoadLibrary(path)
if err != nil {
return nil, err
}
n := &Native{lib: lib}
if err := n.bind(); err != nil {
ffi.FreeLibrary(lib)
return nil, err
}
return n, nil
}
func (n *Native) Close() error {
if n.lib == nil {
return nil
}
ffi.FreeLibrary(n.lib)
n.lib = nil
return nil
}
func (n *Native) RunNetworkInstance(config string) error {
defer pinErrorThread()()
cfg := cString(config)
cfgPtr := unsafe.Pointer(&cfg[0])
var ret int32
err := n.runNetworkInstance.call(unsafe.Pointer(&ret), unsafe.Pointer(&cfgPtr))
runtime.KeepAlive(cfg)
if err != nil {
return err
}
if ret != 0 {
return n.lastError()
}
return nil
}
func (n *Native) CallJSONRPC(serviceName, methodName, domainName, payloadJSON string) (string, error) {
defer pinErrorThread()()
service := cString(serviceName)
method := cString(methodName)
payload := cString(payloadJSON)
servicePtr := unsafe.Pointer(&service[0])
methodPtr := unsafe.Pointer(&method[0])
payloadPtr := unsafe.Pointer(&payload[0])
var domain []byte
var domainPtr unsafe.Pointer
if domainName != "" {
domain = cString(domainName)
domainPtr = unsafe.Pointer(&domain[0])
}
var response unsafe.Pointer
responseArg := unsafe.Pointer(&response)
var ret int32
err := n.callJSONRPC.call(
unsafe.Pointer(&ret),
unsafe.Pointer(&servicePtr),
unsafe.Pointer(&methodPtr),
unsafe.Pointer(&domainPtr),
unsafe.Pointer(&payloadPtr),
unsafe.Pointer(&responseArg),
)
runtime.KeepAlive(service)
runtime.KeepAlive(method)
runtime.KeepAlive(domain)
runtime.KeepAlive(payload)
if err != nil {
return "", err
}
if ret != 0 {
return "", n.lastError()
}
if response == nil {
return "", errors.New("easytier ffi JSON RPC returned nil response")
}
defer func() { _ = n.freeCString(response) }()
return readCString(response), nil
}
func (n *Native) DialContext(ctx context.Context, instance, network, address string) (net.Conn, error) {
if network != "tcp" && network != "tcp4" && network != "tcp6" {
return nil, net.UnknownNetworkError(network)
}
ip, port, err := parseIPPort(address)
if err != nil {
return nil, err
}
timeout := defaultTimeout
if deadline, ok := ctx.Deadline(); ok {
timeout = time.Until(deadline)
}
if timeout <= 0 {
return nil, context.DeadlineExceeded
}
if err := ctx.Err(); err != nil {
return nil, err
}
handle, local, err := n.tcpConnectTo(instance, ip.String(), uint16(port), timeout)
if err != nil {
return nil, err
}
return &Conn{native: n, handle: handle, local: local, remote: &net.TCPAddr{IP: ip, Port: port}}, nil
}
func (n *Native) ListenContext(ctx context.Context, instance, network, address string) (net.Listener, error) {
if network != "tcp" && network != "tcp4" && network != "tcp6" {
return nil, net.UnknownNetworkError(network)
}
port, err := parseListenPort(address)
if err != nil {
return nil, err
}
timeout := defaultTimeout
if deadline, ok := ctx.Deadline(); ok {
timeout = time.Until(deadline)
}
if timeout <= 0 {
return nil, context.DeadlineExceeded
}
if err := ctx.Err(); err != nil {
return nil, err
}
handle, local, err := n.tcpBindTo(instance, uint16(port), timeout)
if err != nil {
return nil, err
}
return &Listener{native: n, handle: handle, addr: local}, nil
}
func (c *Conn) Read(b []byte) (int, error) {
if c.closed.Load() {
return 0, net.ErrClosed
}
n, err := c.native.tcpReadFrom(c.handle, b, c.rd.timeout(defaultTimeout))
if err != nil {
return 0, opError("read", c.remote, err)
}
if n == 0 {
return 0, io.EOF
}
return n, nil
}
func (c *Conn) Write(b []byte) (int, error) {
if c.closed.Load() {
return 0, net.ErrClosed
}
n, err := c.native.tcpWriteTo(c.handle, b, c.wd.timeout(defaultTimeout))
if err != nil {
return 0, opError("write", c.remote, err)
}
return n, nil
}
func (c *Conn) Close() error {
if !c.closed.CompareAndSwap(false, true) {
return net.ErrClosed
}
return c.native.tcpCloseHandle(c.handle)
}
func (c *Conn) LocalAddr() net.Addr { return c.local }
func (c *Conn) RemoteAddr() net.Addr { return c.remote }
func (c *Conn) SetDeadline(t time.Time) error { c.rd.set(t); c.wd.set(t); return nil }
func (c *Conn) SetReadDeadline(t time.Time) error { c.rd.set(t); return nil }
func (c *Conn) SetWriteDeadline(t time.Time) error { c.wd.set(t); return nil }
func (l *Listener) Accept() (net.Conn, error) {
if l.closed.Load() {
return nil, net.ErrClosed
}
for {
handle, local, peer, err := l.native.tcpAcceptFrom(l.handle, defaultTimeout)
if err == nil {
return &Conn{native: l.native, handle: handle, local: local, remote: peer}, nil
}
if l.closed.Load() {
return nil, net.ErrClosed
}
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
continue
}
return nil, opError("accept", l.addr, err)
}
}
func (l *Listener) Close() error {
if !l.closed.CompareAndSwap(false, true) {
return net.ErrClosed
}
return l.native.tcpListenerCloseHandle(l.handle)
}
func (l *Listener) Addr() net.Addr { return l.addr }
func (n *Native) bind() error {
return errors.Join(
n.bindSym(&n.runNetworkInstance, "run_network_instance", types.SInt32TypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.callJSONRPC, "call_json_rpc", types.SInt32TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.getErrorMsg, "get_error_msg", types.VoidTypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.freeString, "free_string", types.VoidTypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.tcpConnect, "data_plane_tcp_connect", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.tcpBind, "data_plane_tcp_bind", types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt16TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.tcpAccept, "data_plane_tcp_accept", types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor, types.PointerTypeDescriptor),
n.bindSym(&n.tcpRead, "data_plane_tcp_read", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor),
n.bindSym(&n.tcpWrite, "data_plane_tcp_write", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor, types.PointerTypeDescriptor, types.UInt32TypeDescriptor, types.UInt64TypeDescriptor),
n.bindSym(&n.tcpClose, "data_plane_tcp_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor),
n.bindSym(&n.tcpListenerClose, "data_plane_tcp_listener_close", types.SInt32TypeDescriptor, types.UInt64TypeDescriptor),
)
}
func (n *Native) bindSym(dst *symCall, name string, ret *types.TypeDescriptor, args ...*types.TypeDescriptor) error {
sym, err := ffi.GetSymbol(n.lib, name)
if err != nil {
return err
}
if err := ffi.PrepareCallInterface(&dst.cif, types.DefaultCall, ret, args); err != nil {
return err
}
dst.fn = sym
return nil
}
func (s *symCall) call(ret unsafe.Pointer, args ...unsafe.Pointer) error {
// `ffi.CallFunction` and libffi `ffi_call` are safe to invoke concurrently
// because `cif` is prepared once during binding and only read afterwards.
return ffi.CallFunction(&s.cif, s.fn, ret, args)
}
func (n *Native) tcpConnectTo(instance, ip string, port uint16, timeout time.Duration) (uint64, *net.TCPAddr, error) {
defer pinErrorThread()()
inst := cString(instance)
dst := cString(ip)
instPtr := unsafe.Pointer(&inst[0])
dstPtr := unsafe.Pointer(&dst[0])
timeoutMS := uint64(timeout / time.Millisecond)
var handle uint64
var outIP unsafe.Pointer
outIPArg := unsafe.Pointer(&outIP)
var outPort uint16
outPortArg := unsafe.Pointer(&outPort)
err := n.tcpConnect.call(
unsafe.Pointer(&handle),
unsafe.Pointer(&instPtr),
unsafe.Pointer(&dstPtr),
unsafe.Pointer(&port),
unsafe.Pointer(&timeoutMS),
unsafe.Pointer(&outIPArg),
unsafe.Pointer(&outPortArg),
)
runtime.KeepAlive(inst)
runtime.KeepAlive(dst)
if err != nil {
return 0, nil, err
}
if handle == 0 {
return 0, nil, n.lastError()
}
return handle, n.takeTCPAddr(outIP, outPort), nil
}
func (n *Native) tcpBindTo(instance string, port uint16, timeout time.Duration) (uint64, *net.TCPAddr, error) {
defer pinErrorThread()()
inst := cString(instance)
instPtr := unsafe.Pointer(&inst[0])
timeoutMS := uint64(timeout / time.Millisecond)
var handle uint64
var outIP unsafe.Pointer
outIPArg := unsafe.Pointer(&outIP)
var outPort uint16
outPortArg := unsafe.Pointer(&outPort)
err := n.tcpBind.call(
unsafe.Pointer(&handle),
unsafe.Pointer(&instPtr),
unsafe.Pointer(&port),
unsafe.Pointer(&timeoutMS),
unsafe.Pointer(&outIPArg),
unsafe.Pointer(&outPortArg),
)
runtime.KeepAlive(inst)
if err != nil {
return 0, nil, err
}
if handle == 0 {
return 0, nil, n.lastError()
}
return handle, n.takeTCPAddr(outIP, outPort), nil
}
func (n *Native) tcpAcceptFrom(handle uint64, timeout time.Duration) (uint64, *net.TCPAddr, *net.TCPAddr, error) {
defer pinErrorThread()()
timeoutMS := uint64(timeout / time.Millisecond)
var stream uint64
var outLocalIP unsafe.Pointer
outLocalIPArg := unsafe.Pointer(&outLocalIP)
var outLocalPort uint16
outLocalPortArg := unsafe.Pointer(&outLocalPort)
var outPeerIP unsafe.Pointer
outPeerIPArg := unsafe.Pointer(&outPeerIP)
var outPeerPort uint16
outPeerPortArg := unsafe.Pointer(&outPeerPort)
err := n.tcpAccept.call(
unsafe.Pointer(&stream),
unsafe.Pointer(&handle),
unsafe.Pointer(&timeoutMS),
unsafe.Pointer(&outLocalIPArg),
unsafe.Pointer(&outLocalPortArg),
unsafe.Pointer(&outPeerIPArg),
unsafe.Pointer(&outPeerPortArg),
)
if err != nil {
return 0, nil, nil, err
}
if stream == 0 {
return 0, nil, nil, n.lastError()
}
return stream, n.takeTCPAddr(outLocalIP, outLocalPort), n.takeTCPAddr(outPeerIP, outPeerPort), nil
}
func (n *Native) tcpReadFrom(handle uint64, buf []byte, timeout time.Duration) (int, error) {
if len(buf) == 0 {
return 0, nil
}
defer pinErrorThread()()
var ret int32
bufPtr := unsafe.Pointer(&buf[0])
length := uint32(len(buf))
timeoutMS := uint64(timeout / time.Millisecond)
err := n.tcpRead.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle), unsafe.Pointer(&bufPtr), unsafe.Pointer(&length), unsafe.Pointer(&timeoutMS))
runtime.KeepAlive(buf)
if err != nil {
return 0, err
}
if ret < 0 {
return 0, n.lastError()
}
return int(ret), nil
}
func (n *Native) tcpWriteTo(handle uint64, buf []byte, timeout time.Duration) (int, error) {
if len(buf) == 0 {
return 0, nil
}
defer pinErrorThread()()
var ret int32
bufPtr := unsafe.Pointer(&buf[0])
length := uint32(len(buf))
timeoutMS := uint64(timeout / time.Millisecond)
err := n.tcpWrite.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle), unsafe.Pointer(&bufPtr), unsafe.Pointer(&length), unsafe.Pointer(&timeoutMS))
runtime.KeepAlive(buf)
if err != nil {
return 0, err
}
if ret < 0 {
return 0, n.lastError()
}
return int(ret), nil
}
func (n *Native) tcpCloseHandle(handle uint64) error {
defer pinErrorThread()()
var ret int32
if err := n.tcpClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil {
return err
}
if ret != 0 {
return n.lastError()
}
return nil
}
func (n *Native) tcpListenerCloseHandle(handle uint64) error {
defer pinErrorThread()()
var ret int32
if err := n.tcpListenerClose.call(unsafe.Pointer(&ret), unsafe.Pointer(&handle)); err != nil {
return err
}
if ret != 0 {
return n.lastError()
}
return nil
}
// pinErrorThread ties an FFI op to the get_error_msg that reads its result: the
// Rust side stores the last error in a thread-local, so the goroutine must not
// migrate to another OS thread between the two calls. Use as `defer pinErrorThread()()`
// at the start of any wrapper that reports failures through lastError.
func pinErrorThread() func() {
runtime.LockOSThread()
return runtime.UnlockOSThread
}
func (n *Native) lastError() error {
var out unsafe.Pointer
outArg := unsafe.Pointer(&out)
if err := n.getErrorMsg.call(nil, unsafe.Pointer(&outArg)); err != nil {
return err
}
if out == nil {
return errors.New("easytier ffi call failed")
}
msg := readCString(out)
_ = n.freeCString(out)
if strings.Contains(msg, "timed out") {
return timeoutError(msg)
}
return errors.New(msg)
}
func (n *Native) freeCString(ptr unsafe.Pointer) error {
if ptr == nil {
return nil
}
return n.freeString.call(nil, unsafe.Pointer(&ptr))
}
func (n *Native) takeTCPAddr(ipPtr unsafe.Pointer, port uint16) *net.TCPAddr {
if ipPtr == nil {
return nil
}
ip := net.ParseIP(readCString(ipPtr))
_ = n.freeCString(ipPtr)
return &net.TCPAddr{IP: ip, Port: int(port)}
}
func (d *atomicDeadline) set(t time.Time) {
if t.IsZero() {
d.v.Store(0)
return
}
d.v.Store(t.UnixNano())
}
func (d *atomicDeadline) timeout(fallback time.Duration) time.Duration {
ns := d.v.Load()
if ns == 0 {
return fallback
}
remaining := time.Until(time.Unix(0, ns))
if remaining <= 0 {
return time.Millisecond
}
return remaining
}
func (e timeoutError) Error() string { return string(e) }
func (e timeoutError) Timeout() bool { return true }
func (e timeoutError) Temporary() bool { return true }
func opError(op string, addr net.Addr, err error) error {
return &net.OpError{Op: op, Net: "easytier", Addr: addr, Err: err}
}
func parseIPPort(address string) (net.IP, int, error) {
host, portStr, err := net.SplitHostPort(address)
if err != nil {
return nil, 0, err
}
ip := net.ParseIP(host)
if ip == nil {
return nil, 0, fmt.Errorf("easytier ffi requires an IP address, got %q", host)
}
port, err := strconv.ParseUint(portStr, 10, 16)
if err != nil {
return nil, 0, err
}
return ip, int(port), nil
}
func parseListenPort(address string) (int, error) {
host, portStr, err := net.SplitHostPort(address)
if err != nil {
return 0, err
}
if host != "" {
ip := net.ParseIP(host)
if ip == nil {
return 0, fmt.Errorf("easytier ffi requires an IP address, got %q", host)
}
if !ip.IsUnspecified() {
return 0, fmt.Errorf("easytier ffi listen address must be unspecified, got %q", host)
}
}
port, err := strconv.ParseUint(portStr, 10, 16)
if err != nil {
return 0, err
}
return int(port), nil
}
func cString(s string) []byte {
if strings.ContainsRune(s, 0) {
panic("easytier ffi string contains NUL")
}
return append([]byte(s), 0)
}
func readCString(ptr unsafe.Pointer) string {
if ptr == nil {
return ""
}
var b []byte
for p := uintptr(ptr); ; p++ {
c := *(*byte)(unsafe.Pointer(p))
if c == 0 {
return string(b)
}
b = append(b, c)
}
}
func defaultLibraryPath() string {
if p := os.Getenv("EASYTIER_FFI_LIB"); p != "" {
return p
}
switch runtime.GOOS {
case "darwin":
return "../../../../target/debug/libeasytier_ffi.dylib"
case "windows":
return "..\\..\\..\\..\\target\\debug\\easytier_ffi.dll"
default:
return "../../../../target/debug/libeasytier_ffi.so"
}
}
var _ net.Conn = (*Conn)(nil)
var _ net.Listener = (*Listener)(nil)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,360 @@
package easytierffi
import (
"context"
"fmt"
"io"
"net"
"os"
"strconv"
"testing"
"time"
)
const asyncLocalTestTimeout = 120 * time.Second
func TestAsyncSymbolBinding(t *testing.T) {
n := openAsyncForTest(t)
status, err := n.opWaitStatus(0, 0)
if err != nil {
t.Fatal(err)
}
if status != dataPlaneOpInvalid {
t.Fatalf("expected invalid status for op 0, got %d", status)
}
}
func TestAsyncLocalTwoNodeTCPAndUDP(t *testing.T) {
n := openAsyncForTest(t)
topology := startLocalAsyncTopology(t, n)
ctx, cancel := context.WithTimeout(context.Background(), asyncLocalTestTimeout)
defer cancel()
runAsyncTCPPingPong(t, ctx, n, topology)
runAsyncUDPPingPong(t, ctx, n, topology)
}
type localAsyncTopology struct {
dialerInstance string
listenerInstance string
listenerIP string
}
func openAsyncForTest(t *testing.T) *AsyncNative {
t.Helper()
libraryPath := defaultLibraryPath()
if _, err := os.Stat(libraryPath); err != nil {
if os.IsNotExist(err) {
t.Skipf("build easytier-ffi with ffi-dataplane before running async tests: %v", err)
}
t.Fatalf("stat async ffi library: %v", err)
}
n, err := OpenAsync(libraryPath)
if err != nil {
t.Fatalf("open async ffi library: %v", err)
}
t.Cleanup(func() {
if err := n.Close(); err != nil {
t.Errorf("close async native: %v", err)
}
})
return n
}
func startLocalAsyncTopology(t *testing.T, n *AsyncNative) localAsyncTopology {
t.Helper()
suffix := strconv.FormatInt(time.Now().UnixNano(), 10)
networkName := "ffi-async-" + suffix
networkSecret := "ffi-async-secret-" + suffix
listenerInstance := "ffi-async-listener-" + suffix
dialerInstance := "ffi-async-dialer-" + suffix
listenerIP := "10.251.1.2"
dialerIP := "10.251.1.1"
listenerPort := freeLocalTCPPort(t)
listenerEndpoint := fmt.Sprintf("tcp://127.0.0.1:%d", listenerPort)
t.Cleanup(func() {
if err := n.deleteNetworkInstances([]string{dialerInstance, listenerInstance}); err != nil {
t.Errorf("cleanup async test EasyTier instances: %v", err)
}
})
listenerConfig := localAsyncConfig(
listenerInstance,
listenerIP,
networkName,
networkSecret,
[]string{listenerEndpoint},
nil,
)
dialerConfig := localAsyncConfig(
dialerInstance,
dialerIP,
networkName,
networkSecret,
nil,
[]string{listenerEndpoint},
)
if err := n.RunNetworkInstance(listenerConfig); err != nil {
t.Fatalf("start listener instance: %v", err)
}
if err := n.RunNetworkInstance(dialerConfig); err != nil {
t.Fatalf("start dialer instance: %v", err)
}
return localAsyncTopology{
dialerInstance: dialerInstance,
listenerInstance: listenerInstance,
listenerIP: listenerIP,
}
}
func localAsyncConfig(instance, ipv4, networkName, networkSecret string, listeners, peers []string) string {
config := fmt.Sprintf(`instance_name = %s
ipv4 = %s
listeners = %s
[network_identity]
network_name = %s
network_secret = %s
[flags]
no_tun = true
bind_device = false
`,
strconv.Quote(instance),
strconv.Quote(ipv4),
tomlStringList(listeners),
strconv.Quote(networkName),
strconv.Quote(networkSecret),
)
for _, peer := range peers {
config += fmt.Sprintf("\n[[peer]]\nuri = %s\n", strconv.Quote(peer))
}
return config
}
func tomlStringList(values []string) string {
if len(values) == 0 {
return "[]"
}
out := "["
for i, value := range values {
if i > 0 {
out += ", "
}
out += strconv.Quote(value)
}
return out + "]"
}
func freeLocalTCPPort(t *testing.T) int {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("allocate local tcp port: %v", err)
}
defer listener.Close()
return listener.Addr().(*net.TCPAddr).Port
}
func runAsyncTCPPingPong(t *testing.T, ctx context.Context, n *AsyncNative, topology localAsyncTopology) {
t.Helper()
listener, listenerAddr := eventuallyTCPListen(t, ctx, n, topology.listenerInstance)
tcpCtx, cancel := context.WithCancel(ctx)
accepted := make(chan error, 1)
defer waitForAsyncHelper(t, accepted, "tcp accept helper")
defer cancel()
defer listener.Close()
go func() {
conn, err := listener.Accept()
if err != nil {
accepted <- fmt.Errorf("accept tcp stream: %w", err)
return
}
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(30 * time.Second))
payload := make([]byte, len("ping"))
if _, err := io.ReadFull(conn, payload); err != nil {
accepted <- fmt.Errorf("read tcp ping: %w", err)
return
}
if string(payload) != "ping" {
accepted <- fmt.Errorf("expected tcp ping, got %q", string(payload))
return
}
if _, err := conn.Write([]byte("pong")); err != nil {
accepted <- fmt.Errorf("write tcp pong: %w", err)
return
}
accepted <- nil
}()
conn, err := eventuallyTCPDial(t, tcpCtx, n, topology.dialerInstance, topology.listenerIP, listenerAddr.Port)
if err != nil {
t.Fatal(err)
}
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(30 * time.Second))
if _, err := conn.Write([]byte("ping")); err != nil {
t.Fatalf("write tcp ping: %v", err)
}
payload := make([]byte, len("pong"))
if _, err := io.ReadFull(conn, payload); err != nil {
t.Fatalf("read tcp pong: %v", err)
}
if string(payload) != "pong" {
t.Fatalf("expected tcp pong, got %q", string(payload))
}
}
func eventuallyTCPListen(t *testing.T, ctx context.Context, n *AsyncNative, instance string) (net.Listener, *net.TCPAddr) {
t.Helper()
var lastErr error
for attempt := 1; ctx.Err() == nil; attempt++ {
attemptCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
listener, err := n.ListenContext(attemptCtx, instance, "tcp", "0.0.0.0:0")
cancel()
if err == nil {
addr := listener.Addr().(*net.TCPAddr)
t.Logf("async tcp bind succeeded on attempt %d at %s", attempt, addr)
return listener, addr
}
lastErr = err
t.Logf("attempt %d: async tcp bind failed: %v", attempt, err)
waitForRetry(ctx, 500*time.Millisecond)
}
t.Fatalf("async tcp bind never succeeded: %v", lastErr)
panic("unreachable")
}
func eventuallyTCPDial(t *testing.T, ctx context.Context, n *AsyncNative, instance, ip string, port int) (net.Conn, error) {
t.Helper()
address := net.JoinHostPort(ip, strconv.Itoa(port))
var lastErr error
for attempt := 1; ctx.Err() == nil; attempt++ {
attemptCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
conn, err := n.DialContext(attemptCtx, instance, "tcp", address)
cancel()
if err == nil {
t.Logf("async tcp connect succeeded on attempt %d to %s", attempt, address)
return conn, nil
}
lastErr = err
t.Logf("attempt %d: async tcp connect failed: %v", attempt, err)
waitForRetry(ctx, 500*time.Millisecond)
}
return nil, fmt.Errorf("async tcp connect never succeeded: %w", lastErr)
}
func runAsyncUDPPingPong(t *testing.T, ctx context.Context, n *AsyncNative, topology localAsyncTopology) {
t.Helper()
dialerSocket, err := n.UDPBindContext(ctx, topology.dialerInstance, 0)
if err != nil {
t.Fatalf("bind dialer udp socket: %v", err)
}
listenerSocket, err := n.UDPBindContext(ctx, topology.listenerInstance, 0)
if err != nil {
t.Fatalf("bind listener udp socket: %v", err)
}
udpCtx, cancel := context.WithCancel(ctx)
warmupDone := make(chan error, 1)
received := make(chan error, 1)
defer waitForAsyncHelper(t, received, "udp receive helper")
defer cancel()
defer listenerSocket.Close()
defer dialerSocket.Close()
go func() {
if _, err := listenerSocket.SendTo(udpCtx, []byte("warmup"), dialerSocket.LocalAddr()); err != nil {
err = fmt.Errorf("send udp warmup: %w", err)
warmupDone <- err
received <- err
return
}
warmupDone <- nil
payload, from, err := listenerSocket.RecvFrom(udpCtx, 512)
if err != nil {
received <- fmt.Errorf("recv udp ping: %w", err)
return
}
if string(payload) != "ping" {
received <- fmt.Errorf("expected udp ping, got %q", string(payload))
return
}
if _, err := listenerSocket.SendTo(udpCtx, []byte("pong"), from); err != nil {
received <- fmt.Errorf("send udp pong: %w", err)
return
}
received <- nil
}()
select {
case err := <-warmupDone:
if err != nil {
t.Fatal(err)
}
case <-udpCtx.Done():
t.Fatal(udpCtx.Err())
}
target := &net.UDPAddr{IP: net.ParseIP(topology.listenerIP), Port: listenerSocket.LocalAddr().Port}
if _, err := dialerSocket.SendTo(udpCtx, []byte("ping"), target); err != nil {
t.Fatalf("send udp ping: %v", err)
}
for {
payload, from, err := dialerSocket.RecvFrom(udpCtx, 512)
if err != nil {
t.Fatalf("recv udp pong: %v", err)
}
if string(payload) == "pong" {
if !from.IP.Equal(target.IP) || from.Port != target.Port {
t.Fatalf("expected udp pong from %s, got %s", target, from)
}
break
}
t.Logf("skipping udp datagram from %s: %q", from, string(payload))
}
}
func waitForAsyncHelper(t *testing.T, done <-chan error, name string) {
t.Helper()
select {
case err := <-done:
if err != nil {
t.Errorf("%s: %v", name, err)
}
case <-time.After(10 * time.Second):
t.Errorf("%s did not stop", name)
}
}
func waitForRetry(ctx context.Context, delay time.Duration) {
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
case <-ctx.Done():
}
}
@@ -0,0 +1,140 @@
package easytierffi
import (
"context"
"fmt"
"io"
"net"
"os"
"strconv"
"strings"
"testing"
"time"
)
func TestSSHIntegration(t *testing.T) {
config := os.Getenv("EASYTIER_FFI_CONFIG")
instance := os.Getenv("EASYTIER_FFI_INSTANCE")
target := os.Getenv("EASYTIER_FFI_TARGET")
if config == "" || instance == "" || target == "" {
t.Skip("set EASYTIER_FFI_CONFIG, EASYTIER_FFI_INSTANCE and EASYTIER_FFI_TARGET to run integration test")
}
n, err := Open(defaultLibraryPath())
if err != nil {
t.Fatal(err)
}
defer n.Close()
if err := n.RunNetworkInstance(config); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
defer cancel()
var lastErr error
for attempt := 1; ctx.Err() == nil; attempt++ {
conn, err := n.DialContext(ctx, instance, "tcp", target)
if err != nil {
lastErr = err
t.Logf("attempt %d: dial failed: %v", attempt, err)
time.Sleep(3 * time.Second)
continue
}
_ = conn.SetReadDeadline(time.Now().Add(10 * time.Second))
buf := make([]byte, 128)
nn, err := conn.Read(buf)
_ = conn.Close()
if err != nil {
lastErr = err
t.Logf("attempt %d: read failed: %v", attempt, err)
time.Sleep(3 * time.Second)
continue
}
banner := string(buf[:nn])
if !strings.HasPrefix(banner, "SSH-") {
t.Fatalf("attempt %d: expected SSH banner, got %q", attempt, banner)
}
t.Logf("attempt %d: got banner %q", attempt, strings.TrimRight(banner, "\r\n"))
return
}
t.Fatalf("never got SSH banner, last err: %v", lastErr)
}
func TestTCPListenIntegration(t *testing.T) {
config := os.Getenv("EASYTIER_FFI_LISTEN_CONFIG")
instance := os.Getenv("EASYTIER_FFI_LISTEN_INSTANCE")
listenPort := os.Getenv("EASYTIER_FFI_LISTEN_PORT")
if config == "" || instance == "" || listenPort == "" {
t.Skip("set EASYTIER_FFI_LISTEN_CONFIG, EASYTIER_FFI_LISTEN_INSTANCE and EASYTIER_FFI_LISTEN_PORT to run integration test")
}
port, err := strconv.ParseUint(listenPort, 10, 16)
if err != nil {
t.Fatal(err)
}
n, err := Open(defaultLibraryPath())
if err != nil {
t.Fatal(err)
}
defer n.Close()
if err := n.RunNetworkInstance(config); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
defer cancel()
// Data-plane readiness is asynchronous: the instance must finish starting
// before the data plane accepts binds. Retry until ready or ctx expires.
var listener net.Listener
for attempt := 1; ; attempt++ {
listener, err = n.ListenContext(ctx, instance, "tcp", net.JoinHostPort("0.0.0.0", strconv.Itoa(int(port))))
if err == nil {
break
}
if ctx.Err() != nil {
t.Fatalf("bind never succeeded, last err: %v", err)
}
t.Logf("attempt %d: bind failed: %v", attempt, err)
time.Sleep(3 * time.Second)
}
t.Logf("listening on %s; connect from another EasyTier peer and send ping", listener.Addr())
accepted := make(chan error, 1)
go func() {
conn, err := listener.Accept()
if err != nil {
accepted <- err
return
}
defer conn.Close()
_ = conn.SetDeadline(time.Now().Add(10 * time.Second))
buf := make([]byte, 4)
if _, err := io.ReadFull(conn, buf); err != nil {
accepted <- err
return
}
if string(buf) != "ping" {
accepted <- fmt.Errorf("expected %q, got %q", "ping", string(buf))
return
}
_, err = conn.Write([]byte("pong"))
accepted <- err
}()
select {
case err := <-accepted:
_ = listener.Close()
if err != nil {
t.Fatal(err)
}
case <-ctx.Done():
_ = listener.Close()
t.Fatal(ctx.Err())
}
}
@@ -0,0 +1,5 @@
module easytierffi-example
go 1.25
require github.com/go-webgpu/goffi v0.4.1
@@ -0,0 +1,2 @@
github.com/go-webgpu/goffi v0.4.1 h1:2hQH5XXloxTyTtIleYv+Rajlwzp6UOETURhSZ5+zJxU=
github.com/go-webgpu/goffi v0.4.1/go.mod h1:wfoxNsJkU+5RFbV1kNN1kunhc1lFHuJKK3zpgx08/uM=
@@ -0,0 +1,575 @@
use std::{
cell::Cell,
collections::HashSet,
ffi::{CString, c_char, c_int, c_void},
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
},
};
use easytier::{
common::{
MachineIdOptions,
config::{ConfigLoader as _, TomlConfigLoader},
},
tunnel::TunnelScheme,
web_client::{WebClient, WebClientHooks, run_web_client},
};
use uuid::Uuid;
use crate::{
data_plane::remove_data_plane_handles_by_instance_ids,
error::set_error_msg,
state::{
ASYNC_RUNTIME, INSTANCE_MANAGER, INSTANCE_MUTATION_LOCK, INSTANCE_NAME_ID_MAP,
lock_remote_instance_mutation, remove_instance_name_ids,
},
strings::{c_str_to_string, optional_c_str_to_string},
types::ConfigServerEventCallback,
};
thread_local! {
static IN_CONFIG_SERVER_CALLBACK: Cell<bool> = const { Cell::new(false) };
}
static CONFIG_SERVER_CLIENT: once_cell::sync::Lazy<Mutex<Option<ManagedConfigServerClient>>> =
once_cell::sync::Lazy::new(|| Mutex::new(None));
static CONFIG_SERVER_CLIENT_ACTIVE: once_cell::sync::Lazy<AtomicBool> =
once_cell::sync::Lazy::new(|| AtomicBool::new(false));
static CONFIG_SERVER_CLIENT_STOPPING: once_cell::sync::Lazy<AtomicBool> =
once_cell::sync::Lazy::new(|| AtomicBool::new(false));
static LAST_CONFIG_SERVER_CALLBACK_ERROR: once_cell::sync::Lazy<Mutex<Option<String>>> =
once_cell::sync::Lazy::new(|| Mutex::new(None));
pub(crate) struct ConfigServerCallbackScope;
impl ConfigServerCallbackScope {
pub(crate) fn enter() -> Self {
IN_CONFIG_SERVER_CALLBACK.with(|in_callback| in_callback.set(true));
Self
}
}
impl Drop for ConfigServerCallbackScope {
fn drop(&mut self) {
IN_CONFIG_SERVER_CALLBACK.with(|in_callback| in_callback.set(false));
}
}
pub fn in_config_server_callback() -> bool {
IN_CONFIG_SERVER_CALLBACK.with(Cell::get)
}
fn config_server_machine_id_options(machine_id: String) -> MachineIdOptions {
MachineIdOptions {
explicit_machine_id: Some(machine_id),
state_dir: None,
}
}
pub fn validate_config_server_client_options(
config_server_url_s: &str,
machine_id: &str,
) -> Result<(), String> {
if machine_id.trim().is_empty() {
return Err("machine_id is empty".to_string());
}
let config_server_url = match url::Url::parse(config_server_url_s) {
Ok(url) => url,
Err(_) => format!(
"udp://config-server.easytier.cn:22020/{}",
config_server_url_s
)
.parse()
.map_err(|err| format!("failed to parse config server URL: {}", err))?,
};
TunnelScheme::try_from(&config_server_url).map_err(|_| {
format!(
"unsupported config server scheme: {}",
config_server_url.scheme()
)
})?;
let token = config_server_url
.path_segments()
.and_then(|mut segments| segments.next_back())
.map(|segment| percent_encoding::percent_decode_str(segment).decode_utf8())
.transpose()
.map_err(|err| format!("failed to decode config server token: {}", err))?
.map(|token| token.to_string())
.unwrap_or_default();
if token.is_empty() {
return Err("empty token".to_string());
}
Ok(())
}
struct ManagedConfigServerClient {
client: WebClient,
hooks: Arc<ManagedConfigServerClientHooks>,
}
pub(crate) struct ManagedConfigServerClientHooks {
pub(crate) instance_ids: Mutex<HashSet<Uuid>>,
callback_delivery: Mutex<()>,
stopping: AtomicBool,
callback: ConfigServerEventCallback,
user_data: usize,
}
impl ManagedConfigServerClientHooks {
pub(crate) fn new(callback: ConfigServerEventCallback, user_data: *mut c_void) -> Self {
Self {
instance_ids: Mutex::new(HashSet::new()),
callback_delivery: Mutex::new(()),
stopping: AtomicBool::new(false),
callback,
user_data: user_data as usize,
}
}
#[cfg(test)]
pub(crate) fn tracked_instance_ids(&self) -> Vec<Uuid> {
self.instance_ids
.lock()
.map(|guard| guard.iter().copied().collect())
.unwrap_or_default()
}
fn remove_tracked_instance_ids(&self, ids: &[Uuid]) -> Result<Vec<Uuid>, String> {
let mut guard = self.instance_ids.lock().map_err(|err| err.to_string())?;
Ok(ids
.iter()
.filter_map(|id| guard.remove(id).then_some(*id))
.collect())
}
fn validate_instance_name(&self, inst_name: &str, inst_id: Uuid) -> Result<(), String> {
if let Some(existing_id) = INSTANCE_NAME_ID_MAP.get(inst_name).map(|id| *id)
&& existing_id != inst_id
{
return Err(format!("instance name {} already exists", inst_name));
}
Ok(())
}
fn commit_instance_name(&self, inst_name: String, inst_id: Uuid) -> Result<(), String> {
INSTANCE_NAME_ID_MAP.retain(|_, existing_id| *existing_id != inst_id);
self.validate_instance_name(&inst_name, inst_id)?;
INSTANCE_NAME_ID_MAP.insert(inst_name, inst_id);
Ok(())
}
pub(crate) fn start_stopping(&self) -> Vec<Uuid> {
let _delivery_guard = if in_config_server_callback() {
None
} else {
self.callback_delivery.lock().ok()
};
let mut guard = match self.instance_ids.lock() {
Ok(guard) => guard,
Err(_) => return Vec::new(),
};
self.stopping.store(true, Ordering::Release);
guard.drain().collect()
}
pub(crate) fn note_callback_error(&self, error: String) {
log::warn!("config server event callback failed: {}", error);
if let Ok(mut guard) = LAST_CONFIG_SERVER_CALLBACK_ERROR.lock() {
*guard = Some(error);
}
}
fn emit_event_with_delivery_locked(
&self,
event: &str,
instance_id: Uuid,
) -> Result<(), String> {
if self.stopping.load(Ordering::Acquire) {
return Ok(());
}
let Some(callback) = self.callback else {
return Ok(());
};
let instance_name = INSTANCE_MANAGER
.get_instance_name(&instance_id)
.unwrap_or_default();
let network_name = INSTANCE_MANAGER
.get_network_name(&instance_id)
.unwrap_or_default();
let event_json = serde_json::json!({
"event": event,
"success": true,
"instance_id": instance_id.to_string(),
"instance_name": instance_name,
"network_name": network_name,
"error": null,
})
.to_string();
let event_json = CString::new(event_json).map_err(|err| err.to_string())?;
let _callback_scope = ConfigServerCallbackScope::enter();
unsafe {
callback(event_json.as_ptr(), self.user_data as *mut c_void);
}
Ok(())
}
fn emit_event(&self, event: &str, instance_id: Uuid) -> Result<(), String> {
let _delivery_guard = self
.callback_delivery
.lock()
.map_err(|err| err.to_string())?;
self.emit_event_with_delivery_locked(event, instance_id)
}
fn wait_for_callback_delivery(&self) {
if in_config_server_callback() {
return;
}
if let Ok(guard) = self.callback_delivery.lock() {
drop(guard);
}
}
}
#[async_trait::async_trait]
impl WebClientHooks for ManagedConfigServerClientHooks {
fn manages_remote_config_instances(&self) -> bool {
true
}
async fn pre_run_network_instance(&self, cfg: &TomlConfigLoader) -> Result<(), String> {
if self.stopping.load(Ordering::Acquire) {
return Err("config server client is stopping".to_string());
}
let inst_name = cfg.get_inst_name();
let inst_id = cfg.get_id();
self.validate_instance_name(&inst_name, inst_id)
}
async fn post_run_network_instance(&self, id: &Uuid) -> Result<(), String> {
let _delivery_guard = self
.callback_delivery
.lock()
.map_err(|err| err.to_string())?;
let Some(inst_name) = INSTANCE_MANAGER.get_instance_name(id) else {
if !self.stopping.load(Ordering::Acquire) {
return Err(format!("instance {} not found after start", id));
}
return Ok(());
};
{
let _mutation_guard = INSTANCE_MUTATION_LOCK
.lock()
.map_err(|err| err.to_string())?;
if INSTANCE_MANAGER.get_instance_name(id).is_none() {
if !self.stopping.load(Ordering::Acquire) {
return Err(format!("instance {} not found after start", id));
}
return Ok(());
}
let should_delete = {
let mut guard = self.instance_ids.lock().map_err(|err| err.to_string())?;
if self.stopping.load(Ordering::Acquire) {
true
} else {
guard.insert(*id);
false
}
};
if should_delete {
if let Err(err) = INSTANCE_MANAGER.delete_network_instance(vec![*id]) {
return Err(err.to_string());
}
remove_instance_name_ids(&[*id]);
return Ok(());
}
if self.stopping.load(Ordering::Acquire) {
self.remove_tracked_instance_ids(&[*id])?;
remove_instance_name_ids(&[*id]);
return Ok(());
}
if let Err(err) = self.commit_instance_name(inst_name.clone(), *id) {
self.remove_tracked_instance_ids(&[*id])?;
if let Err(delete_err) = INSTANCE_MANAGER.delete_network_instance(vec![*id]) {
return Err(format!(
"{}; failed to delete duplicate instance: {}",
err, delete_err
));
}
return Err(err);
}
if self.stopping.load(Ordering::Acquire) {
self.remove_tracked_instance_ids(&[*id])?;
remove_instance_name_ids(&[*id]);
return Ok(());
}
if INSTANCE_MANAGER.get_instance_name(id).is_none() {
self.remove_tracked_instance_ids(&[*id])?;
remove_instance_name_ids(&[*id]);
return Err(format!(
"instance {} was removed before post-run completed",
id
));
}
}
remove_data_plane_handles_by_instance_ids(&[*id]);
if let Err(err) = self.emit_event_with_delivery_locked("run_network_instance", *id) {
self.note_callback_error(err);
}
Ok(())
}
async fn post_remove_network_instances(&self, ids: &[Uuid]) -> Result<(), String> {
let removed_ids = {
let _mutation_guard = INSTANCE_MUTATION_LOCK
.lock()
.map_err(|err| err.to_string())?;
let removed_ids = self.remove_tracked_instance_ids(ids)?;
remove_instance_name_ids(ids);
remove_data_plane_handles_by_instance_ids(&removed_ids);
removed_ids
};
for id in removed_ids {
if let Err(err) = self.emit_event("delete_network_instance", id) {
self.note_callback_error(err);
}
}
Ok(())
}
}
pub(crate) fn remove_config_server_tracked_instance_ids(ids: &[Uuid]) {
if ids.is_empty() {
return;
}
if let Ok(guard) = CONFIG_SERVER_CLIENT.lock()
&& let Some(managed) = guard.as_ref()
&& let Err(err) = managed.hooks.remove_tracked_instance_ids(ids)
{
log::warn!("failed to remove config server tracked ids: {}", err);
}
}
pub(crate) fn wait_for_config_server_delivery() {
let hooks = CONFIG_SERVER_CLIENT
.lock()
.ok()
.and_then(|guard| guard.as_ref().map(|managed| managed.hooks.clone()));
if let Some(hooks) = hooks {
hooks.wait_for_callback_delivery();
}
}
pub(crate) fn last_callback_error() -> Option<String> {
LAST_CONFIG_SERVER_CALLBACK_ERROR
.lock()
.ok()
.and_then(|guard| guard.clone())
}
pub(crate) fn clear_last_callback_error() {
if let Ok(mut guard) = LAST_CONFIG_SERVER_CALLBACK_ERROR.lock() {
*guard = None;
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn is_config_server_active_or_stopping() -> bool {
CONFIG_SERVER_CLIENT_ACTIVE.load(Ordering::Acquire)
|| CONFIG_SERVER_CLIENT_STOPPING.load(Ordering::Acquire)
}
#[cfg(test)]
pub(crate) fn set_active_for_test(active: bool) {
CONFIG_SERVER_CLIENT_ACTIVE.store(active, Ordering::Release);
}
/// # Safety
/// Start the config server client.
///
/// `config_server_url` must be a valid null-terminated UTF-8 string.
/// `hostname` may be null; if non-null it must be a valid null-terminated UTF-8 string.
/// `machine_id` must be a valid null-terminated UTF-8 string.
/// `event_json` passed to `callback` is valid only during that callback invocation.
pub(crate) unsafe fn start_config_server_client(
config_server_url: *const c_char,
hostname: *const c_char,
machine_id: *const c_char,
secure_mode: bool,
callback: ConfigServerEventCallback,
user_data: *mut c_void,
) -> c_int {
if in_config_server_callback() {
set_error_msg("cannot start config server client from config server callback");
return -1;
}
let config_server_url = match unsafe { c_str_to_string(config_server_url, "config_server_url") }
{
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let hostname = match unsafe { optional_c_str_to_string(hostname, "hostname") } {
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let machine_id = match unsafe { c_str_to_string(machine_id, "machine_id") } {
Err(err) => {
set_error_msg(&err);
return -1;
}
Ok(value) => value,
};
if let Err(err) = validate_config_server_client_options(&config_server_url, &machine_id) {
set_error_msg(&err);
return -1;
}
let mut guard = match CONFIG_SERVER_CLIENT.lock() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&format!("failed to lock config server client: {}", err));
return -1;
}
};
if guard.is_some() {
set_error_msg("config server client already exists");
return -1;
}
if CONFIG_SERVER_CLIENT_STOPPING.load(Ordering::Acquire) {
set_error_msg("config server client is stopping");
return -1;
}
clear_last_callback_error();
#[cfg(feature = "ffi-dataplane")]
let data_plane_usage_guard = match crate::data_plane::lock_for_config_server_start() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
CONFIG_SERVER_CLIENT_ACTIVE.store(true, Ordering::Release);
#[cfg(feature = "ffi-dataplane")]
drop(data_plane_usage_guard);
let hooks = Arc::new(ManagedConfigServerClientHooks::new(callback, user_data));
let client = match ASYNC_RUNTIME.block_on(run_web_client(
&config_server_url,
config_server_machine_id_options(machine_id),
hostname,
secure_mode,
INSTANCE_MANAGER.clone(),
Some(hooks.clone()),
)) {
Ok(client) => client,
Err(err) => {
CONFIG_SERVER_CLIENT_ACTIVE.store(false, Ordering::Release);
set_error_msg(&format!("failed to start config server client: {}", err));
return -1;
}
};
*guard = Some(ManagedConfigServerClient { client, hooks });
0
}
pub(crate) fn stop_config_server_client() -> c_int {
if in_config_server_callback() {
set_error_msg("cannot stop config server client from config server callback");
return -1;
}
let mut guard = match CONFIG_SERVER_CLIENT.lock() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&format!("failed to lock config server client: {}", err));
return -1;
}
};
let Some(managed) = guard.as_ref() else {
CONFIG_SERVER_CLIENT_ACTIVE.store(false, Ordering::Release);
return 0;
};
if CONFIG_SERVER_CLIENT_STOPPING.swap(true, Ordering::AcqRel) {
set_error_msg("config server client is stopping");
return -1;
}
let hooks = managed.hooks.clone();
let managed = guard.take().expect("config server client exists");
drop(guard);
let _remote_mutation_guard = lock_remote_instance_mutation();
let tracked_ids = hooks.start_stopping();
drop(managed);
let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() {
Ok(guard) => guard,
Err(err) => {
hooks.wait_for_callback_delivery();
CONFIG_SERVER_CLIENT_ACTIVE.store(false, Ordering::Release);
CONFIG_SERVER_CLIENT_STOPPING.store(false, Ordering::Release);
set_error_msg(&format!("failed to lock instance mutation: {}", err));
return -1;
}
};
let delete_result = INSTANCE_MANAGER.delete_network_instance(tracked_ids.clone());
if delete_result.is_ok() {
remove_instance_name_ids(&tracked_ids);
remove_data_plane_handles_by_instance_ids(&tracked_ids);
}
drop(_mutation_guard);
hooks.wait_for_callback_delivery();
CONFIG_SERVER_CLIENT_ACTIVE.store(false, Ordering::Release);
CONFIG_SERVER_CLIENT_STOPPING.store(false, Ordering::Release);
if let Err(err) = delete_result {
set_error_msg(&format!(
"failed to delete config server instances: {}",
err
));
return -1;
}
0
}
pub(crate) fn is_config_server_client_connected() -> c_int {
CONFIG_SERVER_CLIENT
.lock()
.ok()
.and_then(|guard| guard.as_ref().map(|managed| managed.client.is_connected()))
.map(i32::from)
.unwrap_or(0)
}
@@ -0,0 +1,928 @@
#[cfg(feature = "ffi-dataplane")]
use std::{
future::Future,
net::{IpAddr, SocketAddr},
sync::{
Arc, RwLock,
atomic::{AtomicU64, Ordering},
},
time::Duration,
};
#[cfg(feature = "ffi-dataplane")]
use dashmap::DashMap;
#[cfg(feature = "ffi-dataplane")]
use easytier::launcher::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket};
#[cfg(feature = "ffi-dataplane")]
use tokio::io::{AsyncReadExt, AsyncWriteExt, ReadHalf, WriteHalf};
#[cfg(feature = "ffi-dataplane")]
use tokio_util::sync::CancellationToken;
#[cfg(feature = "ffi-dataplane")]
use uuid::Uuid;
#[cfg(feature = "ffi-dataplane")]
use crate::{
config_server::{in_config_server_callback, is_config_server_active_or_stopping},
error::{free_string, set_error_msg},
state::{INSTANCE_MANAGER, INSTANCE_NAME_ID_MAP},
};
#[cfg(feature = "ffi-dataplane")]
static NEXT_DATA_PLANE_HANDLE: AtomicU64 = AtomicU64::new(1);
#[cfg(feature = "ffi-dataplane")]
static DATA_PLANE_HANDLES: once_cell::sync::Lazy<DashMap<u64, DataPlaneHandle>> =
once_cell::sync::Lazy::new(DashMap::new);
#[cfg(feature = "ffi-dataplane")]
static DATA_PLANE_USAGE_LOCK: once_cell::sync::Lazy<RwLock<()>> =
once_cell::sync::Lazy::new(|| RwLock::new(()));
#[cfg(feature = "ffi-dataplane")]
pub(crate) struct DataPlaneHandle {
pub(crate) instance_id: uuid::Uuid,
pub(crate) runtime: tokio::runtime::Handle,
// Cancelled by close() to wake any in-flight op on this handle.
pub(crate) close_token: CancellationToken,
pub(crate) resource: DataPlaneResource,
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) struct TcpHalves {
pub(crate) read: tokio::sync::Mutex<ReadHalf<DataPlaneTcpStream>>,
pub(crate) write: tokio::sync::Mutex<WriteHalf<DataPlaneTcpStream>>,
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) enum DataPlaneResource {
Tcp(Arc<TcpHalves>),
TcpListener(Arc<tokio::sync::Mutex<DataPlaneTcpListener>>),
Udp(Arc<DataPlaneUdpSocket>),
}
// Several helper functions for FFI data plane operations to facilitate logic reuse.
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn next_handle() -> u64 {
NEXT_DATA_PLANE_HANDLE.fetch_add(1, Ordering::Relaxed)
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn timeout_duration(timeout_ms: u64) -> Duration {
Duration::from_millis(timeout_ms)
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn cstr_to_string(ptr: *const std::ffi::c_char, name: &str) -> Option<String> {
if ptr.is_null() {
set_error_msg(&format!("{} is null", name));
return None;
}
Some(
unsafe { std::ffi::CStr::from_ptr(ptr) }
.to_string_lossy()
.into_owned(),
)
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_instance_id(inst_name: &str) -> Option<uuid::Uuid> {
INSTANCE_NAME_ID_MAP.get(inst_name).map(|id| *id.value())
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn parse_socket_addr(host: &str, port: u16) -> Option<SocketAddr> {
let ip = match host.parse::<IpAddr>() {
Ok(ip) => ip,
Err(e) => {
set_error_msg(&format!("failed to parse ip address: {}", e));
return None;
}
};
Some(SocketAddr::new(ip, port))
}
/// Encode an IP address for FFI return. Returns `*mut c_char` to match
/// `CString::into_raw`; caller releases it via `free_string`.
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn into_ffi_ip_cstring(ip: IpAddr) -> Option<*mut std::ffi::c_char> {
match std::ffi::CString::new(ip.to_string()) {
Ok(s) => Some(s.into_raw()),
Err(e) => {
set_error_msg(&format!("failed to encode ip: {}", e));
None
}
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_runtime_handle(
inst_id: &uuid::Uuid,
deadline: std::time::Instant,
) -> Option<tokio::runtime::Handle> {
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
let Some(rt) = INSTANCE_MANAGER.data_plane_wait_runtime_handle(inst_id, remaining) else {
set_error_msg("instance runtime is not ready");
return None;
};
Some(rt)
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn insert_tcp_stream_handle(
instance_id: uuid::Uuid,
runtime: tokio::runtime::Handle,
stream: DataPlaneTcpStream,
) -> u64 {
let (rd, wr) = tokio::io::split(stream);
let handle = next_handle();
DATA_PLANE_HANDLES.insert(
handle,
DataPlaneHandle {
instance_id,
runtime,
close_token: CancellationToken::new(),
resource: DataPlaneResource::Tcp(Arc::new(TcpHalves {
read: tokio::sync::Mutex::new(rd),
write: tokio::sync::Mutex::new(wr),
})),
},
);
handle
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn insert_tcp_listener_handle(
instance_id: uuid::Uuid,
runtime: tokio::runtime::Handle,
listener: DataPlaneTcpListener,
) -> u64 {
let handle = next_handle();
DATA_PLANE_HANDLES.insert(
handle,
DataPlaneHandle {
instance_id,
runtime,
close_token: CancellationToken::new(),
resource: DataPlaneResource::TcpListener(Arc::new(tokio::sync::Mutex::new(listener))),
},
);
handle
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn insert_udp_socket_handle(
instance_id: uuid::Uuid,
runtime: tokio::runtime::Handle,
socket: DataPlaneUdpSocket,
) -> u64 {
let handle = next_handle();
DATA_PLANE_HANDLES.insert(
handle,
DataPlaneHandle {
instance_id,
runtime,
close_token: CancellationToken::new(),
resource: DataPlaneResource::Udp(Arc::new(socket)),
},
);
handle
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_tcp_stream(
handle: u64,
) -> Option<(Arc<TcpHalves>, tokio::runtime::Handle, CancellationToken)> {
get_tcp_stream_with_instance(handle)
.map(|(halves, runtime, close_token, _)| (halves, runtime, close_token))
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_tcp_stream_with_instance(
handle: u64,
) -> Option<(
Arc<TcpHalves>,
tokio::runtime::Handle,
CancellationToken,
uuid::Uuid,
)> {
let Some(h) = DATA_PLANE_HANDLES.get(&handle) else {
set_error_msg("tcp stream handle not found");
return None;
};
match &h.resource {
DataPlaneResource::Tcp(halves) => Some((
halves.clone(),
h.runtime.clone(),
h.close_token.clone(),
h.instance_id,
)),
DataPlaneResource::TcpListener(_) | DataPlaneResource::Udp(_) => {
set_error_msg("handle is not a tcp stream");
None
}
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_tcp_listener(
handle: u64,
) -> Option<(
Arc<tokio::sync::Mutex<DataPlaneTcpListener>>,
tokio::runtime::Handle,
CancellationToken,
uuid::Uuid,
)> {
let Some(h) = DATA_PLANE_HANDLES.get(&handle) else {
set_error_msg("tcp listener handle not found");
return None;
};
match &h.resource {
DataPlaneResource::TcpListener(listener) => Some((
listener.clone(),
h.runtime.clone(),
h.close_token.clone(),
h.instance_id,
)),
DataPlaneResource::Tcp(_) | DataPlaneResource::Udp(_) => {
set_error_msg("handle is not a tcp listener");
None
}
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_udp_socket(
handle: u64,
) -> Option<(
Arc<DataPlaneUdpSocket>,
tokio::runtime::Handle,
CancellationToken,
)> {
get_udp_socket_with_instance(handle)
.map(|(socket, runtime, close_token, _)| (socket, runtime, close_token))
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn get_udp_socket_with_instance(
handle: u64,
) -> Option<(
Arc<DataPlaneUdpSocket>,
tokio::runtime::Handle,
CancellationToken,
uuid::Uuid,
)> {
let Some(h) = DATA_PLANE_HANDLES.get(&handle) else {
set_error_msg("udp socket handle not found");
return None;
};
match &h.resource {
DataPlaneResource::Udp(socket) => Some((
socket.clone(),
h.runtime.clone(),
h.close_token.clone(),
h.instance_id,
)),
DataPlaneResource::Tcp(_) | DataPlaneResource::TcpListener(_) => {
set_error_msg("handle is not a udp socket");
None
}
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn remove_data_plane_handles_by_instance_ids(ids: &[Uuid]) {
if ids.is_empty() {
return;
}
let _data_plane_usage_guard = DATA_PLANE_USAGE_LOCK
.write()
.unwrap_or_else(|err| err.into_inner());
DATA_PLANE_HANDLES.retain(|_, handle| {
if ids.contains(&handle.instance_id) {
handle.close_token.cancel();
false
} else {
true
}
});
crate::data_plane_async::remove_ops_by_instance_ids(ids);
}
#[cfg(not(feature = "ffi-dataplane"))]
pub(crate) fn remove_data_plane_handles_by_instance_ids(_ids: &[uuid::Uuid]) {}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn data_plane_rejected() -> bool {
if in_config_server_callback() {
set_error_msg("cannot use data plane from config server callback");
true
} else if is_config_server_active_or_stopping() {
set_error_msg("cannot use data plane while config server client is active");
true
} else {
false
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn enter_data_plane_operation() -> Option<std::sync::RwLockReadGuard<'static, ()>> {
if data_plane_rejected() {
return None;
}
let guard = match DATA_PLANE_USAGE_LOCK.read() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&format!("failed to lock data plane usage: {}", err));
return None;
}
};
if data_plane_rejected() {
return None;
}
Some(guard)
}
/// Run an IO op on the resource's owning runtime, supporting
/// timeout and cancellation.
#[cfg(feature = "ffi-dataplane")]
async fn run_with_cancel<T, F>(
close_token: &CancellationToken,
timeout_ms: u64,
error_prefix: &str,
op: F,
) -> Option<Result<T, std::io::Error>>
where
F: Future<Output = Result<T, std::io::Error>>,
{
tokio::select! {
biased;
_ = close_token.cancelled() => {
set_error_msg(&format!("{}: handle closed", error_prefix));
None
}
res = tokio::time::timeout(timeout_duration(timeout_ms), op) => match res {
Ok(r) => Some(r),
Err(_) => {
set_error_msg(&format!("{} timed out", error_prefix));
None
}
}
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn lock_for_config_server_start()
-> Result<std::sync::RwLockWriteGuard<'static, ()>, String> {
let guard = DATA_PLANE_USAGE_LOCK
.write()
.map_err(|err| format!("failed to lock data plane usage: {}", err))?;
if !DATA_PLANE_HANDLES.is_empty() || crate::data_plane_async::has_live_ops() {
return Err("cannot start config server client while data plane is in use".to_string());
}
Ok(guard)
}
/// # Safety
/// Open a TCP stream through an EasyTier instance data plane. Returns 0 on
/// failure. On success, writes the local socket address chosen for this
/// connection into `out_local_ip` (a heap-allocated C string the caller must
/// release via `free_string`) and `out_local_port`. Both out pointers must be
/// non-null.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_tcp_connect(
inst_name: *const std::ffi::c_char,
dst_ip: *const std::ffi::c_char,
dst_port: std::ffi::c_ushort,
timeout_ms: u64,
out_local_ip: *mut *const std::ffi::c_char,
out_local_port: *mut std::ffi::c_ushort,
) -> u64 {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return 0,
};
if out_local_ip.is_null() || out_local_port.is_null() {
set_error_msg("output pointer is null");
return 0;
}
let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else {
return 0;
};
let Some(dst_ip) = (unsafe { cstr_to_string(dst_ip, "dst_ip") }) else {
return 0;
};
let Some(inst_id) = get_instance_id(&inst_name) else {
set_error_msg("instance not found");
return 0;
};
let Some(dst_addr) = parse_socket_addr(&dst_ip, dst_port) else {
return 0;
};
let deadline = std::time::Instant::now() + timeout_duration(timeout_ms);
let Some(runtime) = get_runtime_handle(&inst_id, deadline) else {
return 0;
};
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
let result =
runtime.block_on(INSTANCE_MANAGER.data_plane_tcp_connect(&inst_id, dst_addr, remaining));
match result {
Ok(stream) => {
let local_addr = stream.local_addr();
let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else {
return 0;
};
let handle = insert_tcp_stream_handle(inst_id, runtime, stream);
unsafe {
*out_local_ip = local_ip as *const std::ffi::c_char;
*out_local_port = local_addr.port();
}
handle
}
Err(e) => {
set_error_msg(&format!("failed to connect tcp data plane: {}", e));
0
}
}
}
/// # Safety
/// Bind a TCP listener through an EasyTier instance data plane. Returns 0 on
/// failure. The local address actually bound is written into `out_local_ip` /
/// `out_local_port`; the caller must release `*out_local_ip` via `free_string`.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_tcp_bind(
inst_name: *const std::ffi::c_char,
local_port: std::ffi::c_ushort,
timeout_ms: u64,
out_local_ip: *mut *const std::ffi::c_char,
out_local_port: *mut std::ffi::c_ushort,
) -> u64 {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return 0,
};
if out_local_ip.is_null() || out_local_port.is_null() {
set_error_msg("output pointer is null");
return 0;
}
let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else {
return 0;
};
let Some(inst_id) = get_instance_id(&inst_name) else {
set_error_msg("instance not found");
return 0;
};
let deadline = std::time::Instant::now() + timeout_duration(timeout_ms);
let Some(runtime) = get_runtime_handle(&inst_id, deadline) else {
return 0;
};
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
let result =
runtime.block_on(INSTANCE_MANAGER.data_plane_tcp_bind(&inst_id, local_port, remaining));
match result {
Ok(listener) => {
let local_addr = listener.local_addr();
let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else {
return 0;
};
let handle = insert_tcp_listener_handle(inst_id, runtime, listener);
unsafe {
*out_local_ip = local_ip as *const std::ffi::c_char;
*out_local_port = local_addr.port();
}
handle
}
Err(e) => {
set_error_msg(&format!("failed to bind tcp data plane: {}", e));
0
}
}
}
/// # Safety
/// Accept one connection from a TCP data-plane listener. Returns a TCP stream
/// handle, or 0 on failure. Local and peer addresses are written into out
/// parameters; returned IP strings must be released via `free_string`.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_tcp_accept(
handle: u64,
timeout_ms: u64,
out_local_ip: *mut *const std::ffi::c_char,
out_local_port: *mut std::ffi::c_ushort,
out_peer_ip: *mut *const std::ffi::c_char,
out_peer_port: *mut std::ffi::c_ushort,
) -> u64 {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return 0,
};
if out_local_ip.is_null()
|| out_local_port.is_null()
|| out_peer_ip.is_null()
|| out_peer_port.is_null()
{
set_error_msg("output pointer is null");
return 0;
}
let Some((listener, runtime, close_token, instance_id)) = get_tcp_listener(handle) else {
return 0;
};
let ret = runtime.block_on(async move {
let mut listener = listener.lock().await;
run_with_cancel(
&close_token,
timeout_ms,
"tcp data plane accept",
listener.accept(),
)
.await
});
match ret {
Some(Ok((stream, peer_addr))) => {
let local_addr = stream.local_addr();
let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else {
return 0;
};
let Some(peer_ip) = into_ffi_ip_cstring(peer_addr.ip()) else {
free_string(local_ip);
return 0;
};
let stream_handle = insert_tcp_stream_handle(instance_id, runtime, stream);
unsafe {
*out_local_ip = local_ip as *const std::ffi::c_char;
*out_local_port = local_addr.port();
*out_peer_ip = peer_ip as *const std::ffi::c_char;
*out_peer_port = peer_addr.port();
}
stream_handle
}
Some(Err(e)) => {
set_error_msg(&format!("failed to accept tcp data plane: {}", e));
0
}
None => 0,
}
}
/// # Safety
/// Read from a TCP data-plane stream.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_tcp_read(
handle: u64,
buf: *mut std::ffi::c_uchar,
len: u32,
timeout_ms: u64,
) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
if buf.is_null() {
set_error_msg("buf is null");
return -1;
}
let Some((halves, runtime, close_token)) = get_tcp_stream(handle) else {
return -1;
};
// Safety: caller-owned buffer outlives this blocking call.
let buf = unsafe { std::slice::from_raw_parts_mut(buf, len as usize) };
runtime.block_on(async move {
let mut rd = halves.read.lock().await;
match run_with_cancel(
&close_token,
timeout_ms,
"failed to read tcp data plane",
rd.read(buf),
)
.await
{
Some(Ok(n)) => n as std::ffi::c_int,
Some(Err(e)) => {
set_error_msg(&format!("failed to read tcp data plane: {}", e));
-1
}
None => -1,
}
})
}
/// # Safety
/// Write to a TCP data-plane stream.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_tcp_write(
handle: u64,
buf: *const std::ffi::c_uchar,
len: u32,
timeout_ms: u64,
) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
if buf.is_null() {
set_error_msg("buf is null");
return -1;
}
let Some((halves, runtime, close_token)) = get_tcp_stream(handle) else {
return -1;
};
let total = len as usize;
// Safety: caller-owned buffer outlives this blocking call.
let buf = unsafe { std::slice::from_raw_parts(buf, total) };
runtime.block_on(async move {
let mut wr = halves.write.lock().await;
// Use `write_all` to honor `net.Conn::Write` semantics on the Go side
// (must write everything or return an error); single `write()` can
// silently short-write and corrupt streams that the caller assumes are
// fully written.
match run_with_cancel(
&close_token,
timeout_ms,
"failed to write tcp data plane",
wr.write_all(buf),
)
.await
{
Some(Ok(())) => total as std::ffi::c_int,
Some(Err(e)) => {
set_error_msg(&format!("failed to write tcp data plane: {}", e));
-1
}
None => -1,
}
})
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn data_plane_tcp_close(handle: u64) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
crate::data_plane_async::cancel_ops_for_handle(handle);
let Some((_, h)) = DATA_PLANE_HANDLES.remove_if(&handle, |_, e| {
matches!(e.resource, DataPlaneResource::Tcp(_))
}) else {
set_error_msg(if DATA_PLANE_HANDLES.contains_key(&handle) {
"handle is not a tcp stream"
} else {
"tcp stream handle not found"
});
return -1;
};
h.close_token.cancel();
if let DataPlaneResource::Tcp(halves) = h.resource {
// Best-effort half-close; if write half is in use, the in-flight call
// observes the cancel token and releases the lock shortly after.
h.runtime.spawn(async move {
if let Ok(mut wr) = halves.write.try_lock() {
let _ = wr.shutdown().await;
}
});
}
0
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn data_plane_tcp_listener_close(handle: u64) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
crate::data_plane_async::cancel_ops_for_handle(handle);
let Some((_, h)) = DATA_PLANE_HANDLES.remove_if(&handle, |_, e| {
matches!(e.resource, DataPlaneResource::TcpListener(_))
}) else {
set_error_msg(if DATA_PLANE_HANDLES.contains_key(&handle) {
"handle is not a tcp listener"
} else {
"tcp listener handle not found"
});
return -1;
};
h.close_token.cancel();
0
}
/// # Safety
/// Bind a UDP socket through an EasyTier instance data plane. Returns 0 on
/// failure. The local address actually bound (which may differ from the
/// requested port when `local_port == 0`) is written into `out_local_ip` /
/// `out_local_port`; the caller must release `*out_local_ip` via `free_string`.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_udp_bind(
inst_name: *const std::ffi::c_char,
local_port: std::ffi::c_ushort,
timeout_ms: u64,
out_local_ip: *mut *const std::ffi::c_char,
out_local_port: *mut std::ffi::c_ushort,
) -> u64 {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return 0,
};
if out_local_ip.is_null() || out_local_port.is_null() {
set_error_msg("output pointer is null");
return 0;
}
let Some(inst_name) = (unsafe { cstr_to_string(inst_name, "inst_name") }) else {
return 0;
};
let Some(inst_id) = get_instance_id(&inst_name) else {
set_error_msg("instance not found");
return 0;
};
let deadline = std::time::Instant::now() + timeout_duration(timeout_ms);
let Some(runtime) = get_runtime_handle(&inst_id, deadline) else {
return 0;
};
let remaining = deadline.saturating_duration_since(std::time::Instant::now());
let result =
runtime.block_on(INSTANCE_MANAGER.data_plane_udp_bind(&inst_id, local_port, remaining));
match result {
Ok(socket) => {
let local_addr = socket.local_addr();
let Some(local_ip) = into_ffi_ip_cstring(local_addr.ip()) else {
return 0;
};
let handle = insert_udp_socket_handle(inst_id, runtime, socket);
unsafe {
*out_local_ip = local_ip as *const std::ffi::c_char;
*out_local_port = local_addr.port();
}
handle
}
Err(e) => {
set_error_msg(&format!("failed to bind udp data plane: {}", e));
0
}
}
}
/// # Safety
/// Send a datagram through a UDP data-plane socket.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_udp_send_to(
handle: u64,
dst_ip: *const std::ffi::c_char,
dst_port: std::ffi::c_ushort,
buf: *const std::ffi::c_uchar,
len: u32,
timeout_ms: u64,
) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
if buf.is_null() {
set_error_msg("buf is null");
return -1;
}
let Some(dst_ip) = (unsafe { cstr_to_string(dst_ip, "dst_ip") }) else {
return -1;
};
let Some(dst_addr) = parse_socket_addr(&dst_ip, dst_port) else {
return -1;
};
let Some((socket, runtime, close_token)) = get_udp_socket(handle) else {
return -1;
};
let total = len as usize;
// Safety: caller-owned buffer outlives this blocking call.
let buf = unsafe { std::slice::from_raw_parts(buf, total) };
runtime.block_on(async move {
match run_with_cancel(
&close_token,
timeout_ms,
"failed to send udp data plane",
socket.send_to(buf, dst_addr),
)
.await
{
Some(Ok(n)) => n as std::ffi::c_int,
Some(Err(e)) => {
set_error_msg(&format!("failed to send udp data plane: {}", e));
-1
}
None => -1,
}
})
}
/// # Safety
/// Receive a datagram from a UDP data-plane socket.
#[cfg(feature = "ffi-dataplane")]
pub(crate) unsafe fn data_plane_udp_recv_from(
handle: u64,
buf: *mut std::ffi::c_uchar,
len: u32,
out_ip: *mut *const std::ffi::c_char,
out_port: *mut std::ffi::c_ushort,
timeout_ms: u64,
) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
if buf.is_null() || out_ip.is_null() || out_port.is_null() {
set_error_msg("output pointer is null");
return -1;
}
let Some((socket, runtime, close_token)) = get_udp_socket(handle) else {
return -1;
};
let total = len as usize;
// Safety: caller-owned buffer outlives this blocking call.
let buf = unsafe { std::slice::from_raw_parts_mut(buf, total) };
let ret = runtime.block_on(run_with_cancel(
&close_token,
timeout_ms,
"udp data plane receive",
socket.recv_from(buf),
));
match ret {
Some(Ok((n, addr))) => {
// The returned ip pointer must be released by the caller via
// `free_string` (which calls `CString::from_raw`, matching
// `CString::into_raw` here).
let Some(ip_cstr) = into_ffi_ip_cstring(addr.ip()) else {
return -1;
};
unsafe {
*out_ip = ip_cstr as *const std::ffi::c_char;
*out_port = addr.port() as std::ffi::c_ushort;
}
n as std::ffi::c_int
}
Some(Err(e)) => {
set_error_msg(&format!("failed to receive udp data plane: {}", e));
-1
}
None => -1,
}
}
#[cfg(feature = "ffi-dataplane")]
pub(crate) fn data_plane_udp_close(handle: u64) -> std::ffi::c_int {
let _data_plane_usage_guard = match enter_data_plane_operation() {
Some(guard) => guard,
None => return -1,
};
crate::data_plane_async::cancel_ops_for_handle(handle);
let Some((_, h)) = DATA_PLANE_HANDLES.remove_if(&handle, |_, e| {
matches!(e.resource, DataPlaneResource::Udp(_))
}) else {
set_error_msg(if DATA_PLANE_HANDLES.contains_key(&handle) {
"handle is not a udp socket"
} else {
"udp socket handle not found"
});
return -1;
};
h.close_token.cancel();
0
}
#[cfg(all(test, feature = "ffi-dataplane"))]
mod tests {
use super::*;
use std::{sync::mpsc, time::Duration};
#[test]
fn config_server_start_waits_for_data_plane_operation() {
let read_guard = DATA_PLANE_USAGE_LOCK.read().unwrap();
let (done_tx, done_rx) = mpsc::channel();
let waiter = std::thread::spawn(move || {
let _write_guard = lock_for_config_server_start().unwrap();
done_tx.send(()).unwrap();
});
assert!(done_rx.recv_timeout(Duration::from_millis(100)).is_err());
drop(read_guard);
done_rx.recv_timeout(Duration::from_secs(5)).unwrap();
waiter.join().unwrap();
}
#[test]
fn instance_cleanup_waits_for_data_plane_operation() {
let read_guard = DATA_PLANE_USAGE_LOCK.read().unwrap();
let instance_id = Uuid::new_v4();
let (done_tx, done_rx) = mpsc::channel();
let cleaner = std::thread::spawn(move || {
remove_data_plane_handles_by_instance_ids(&[instance_id]);
done_tx.send(()).unwrap();
});
assert!(done_rx.recv_timeout(Duration::from_millis(100)).is_err());
drop(read_guard);
done_rx.recv_timeout(Duration::from_secs(5)).unwrap();
cleaner.join().unwrap();
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,65 @@
use std::{
cell::RefCell,
ffi::{CString, c_char},
};
thread_local! {
// # Thread Safety
// set_error_msg and get_error_msg must be called on the same thread to
// get correct error. And since `Handle::block_on` polls the top-level
// future on the calling thread, set_error_msg always runs on the same
// thread as the corresponding get_error_msg.
static ERROR_MSG: RefCell<Vec<u8>> = const { RefCell::new(Vec::new()) };
}
pub(crate) fn set_error_msg(msg: &str) {
ERROR_MSG.with(|cell| {
let mut buf = cell.borrow_mut();
buf.clear();
buf.extend_from_slice(msg.as_bytes());
});
}
fn thread_local_error_msg() -> Option<String> {
ERROR_MSG.with(|cell| {
let buf = cell.borrow();
if buf.is_empty() {
None
} else {
Some(String::from_utf8_lossy(&buf).into_owned())
}
})
}
pub(crate) unsafe fn get_error_msg(out: *mut *const c_char) {
let msg = match (
thread_local_error_msg(),
crate::config_server::last_callback_error(),
) {
(Some(error), Some(callback_error)) => Some(format!(
"{}; config server callback error: {}",
error, callback_error
)),
(Some(error), None) => Some(error),
(None, Some(callback_error)) => {
Some(format!("config server callback error: {}", callback_error))
}
(None, None) => None,
};
let cstr = msg.and_then(|msg| CString::new(msg).ok());
unsafe {
*out = match cstr {
Some(s) => s.into_raw() as *const c_char,
None => std::ptr::null(),
};
}
}
pub(crate) fn free_string(s: *const c_char) {
if s.is_null() {
return;
}
unsafe {
let _ = CString::from_raw(s as *mut c_char);
}
}
@@ -0,0 +1,366 @@
use std::ffi::{CString, c_char, c_int};
use easytier::common::config::{ConfigFileControl, ConfigLoader as _, TomlConfigLoader};
use crate::{
config_server::{
in_config_server_callback, remove_config_server_tracked_instance_ids,
wait_for_config_server_delivery,
},
data_plane::remove_data_plane_handles_by_instance_ids,
error::set_error_msg,
state::{
INSTANCE_MANAGER, INSTANCE_MUTATION_LOCK, INSTANCE_NAME_ID_MAP, instance_name_exists,
lock_remote_instance_mutation,
},
types::KeyValuePair,
};
/// # Safety
/// Set the tun fd
pub(crate) unsafe fn set_tun_fd(inst_name: *const c_char, fd: c_int) -> c_int {
let inst_name = unsafe {
assert!(!inst_name.is_null());
std::ffi::CStr::from_ptr(inst_name)
.to_string_lossy()
.into_owned()
};
if !INSTANCE_NAME_ID_MAP.contains_key(&inst_name) {
return -1;
}
let inst_id = *INSTANCE_NAME_ID_MAP
.get(&inst_name)
.as_ref()
.unwrap()
.value();
match INSTANCE_MANAGER.set_tun_fd(&inst_id, fd) {
Ok(_) => 0,
Err(_) => -1,
}
}
/// # Safety
/// Parse the config
pub(crate) unsafe fn parse_config(cfg_str: *const std::ffi::c_char) -> std::ffi::c_int {
let cfg_str = unsafe {
assert!(!cfg_str.is_null());
std::ffi::CStr::from_ptr(cfg_str)
.to_string_lossy()
.into_owned()
};
if let Err(e) = TomlConfigLoader::new_from_str(&cfg_str) {
set_error_msg(&format!("failed to parse config: {:?}", e));
return -1;
}
0
}
/// # Safety
/// Run the network instance
pub(crate) unsafe fn run_network_instance(cfg_str: *const std::ffi::c_char) -> std::ffi::c_int {
if in_config_server_callback() {
set_error_msg("cannot run network instance from config server callback");
return -1;
}
let cfg_str = unsafe {
assert!(!cfg_str.is_null());
std::ffi::CStr::from_ptr(cfg_str)
.to_string_lossy()
.into_owned()
};
let cfg = match TomlConfigLoader::new_from_str(&cfg_str) {
Ok(cfg) => cfg,
Err(e) => {
set_error_msg(&format!("failed to parse config: {}", e));
return -1;
}
};
let inst_name = cfg.get_inst_name();
wait_for_config_server_delivery();
let _remote_mutation_guard = lock_remote_instance_mutation();
let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&format!("failed to lock instance mutation: {}", err));
return -1;
}
};
if instance_name_exists(&inst_name) {
set_error_msg("instance already exists");
return -1;
}
let instance_id =
match INSTANCE_MANAGER.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) {
Ok(id) => id,
Err(e) => {
set_error_msg(&format!("failed to start instance: {}", e));
return -1;
}
};
INSTANCE_NAME_ID_MAP.insert(inst_name, instance_id);
0
}
unsafe fn parse_instance_names(
inst_names: *const *const c_char,
length: usize,
) -> Option<Vec<String>> {
if length == 0 {
return Some(Vec::new());
}
if inst_names.is_null() {
set_error_msg("inst_names is null");
return None;
}
let names = unsafe { std::slice::from_raw_parts(inst_names, length) };
let mut parsed = Vec::with_capacity(length);
for (index, &name) in names.iter().enumerate() {
if name.is_null() {
set_error_msg(&format!("inst_names[{}] is null", index));
return None;
}
parsed.push(
unsafe { std::ffi::CStr::from_ptr(name) }
.to_string_lossy()
.into_owned(),
);
}
Some(parsed)
}
/// # Safety
/// Retain the network instance
pub(crate) unsafe fn retain_network_instance(
inst_names: *const *const std::ffi::c_char,
length: usize,
) -> std::ffi::c_int {
if in_config_server_callback() {
set_error_msg("cannot retain network instances from config server callback");
return -1;
}
wait_for_config_server_delivery();
let _remote_mutation_guard = lock_remote_instance_mutation();
let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&format!("failed to lock instance mutation: {}", err));
return -1;
}
};
if length == 0 {
let removed_ids = INSTANCE_MANAGER.list_network_instance_ids();
if let Err(e) = INSTANCE_MANAGER.delete_network_instance(removed_ids.clone()) {
set_error_msg(&format!("failed to delete instances: {}", e));
return -1;
}
remove_config_server_tracked_instance_ids(&removed_ids);
remove_data_plane_handles_by_instance_ids(&removed_ids);
INSTANCE_NAME_ID_MAP.clear();
return 0;
}
let Some(inst_names) = (unsafe { parse_instance_names(inst_names, length) }) else {
return -1;
};
let removed_ids = INSTANCE_MANAGER
.list_network_instance_ids()
.into_iter()
.filter(|id| {
INSTANCE_MANAGER
.get_instance_name(id)
.is_none_or(|name| !inst_names.contains(&name))
})
.collect::<Vec<_>>();
if let Err(e) = INSTANCE_MANAGER.delete_network_instance(removed_ids.clone()) {
set_error_msg(&format!("failed to delete instances: {}", e));
return -1;
}
remove_config_server_tracked_instance_ids(&removed_ids);
remove_data_plane_handles_by_instance_ids(&removed_ids);
INSTANCE_NAME_ID_MAP.retain(|k, _| inst_names.contains(k));
0
}
/// # Safety
/// Delete named network instances.
pub(crate) unsafe fn delete_network_instance(
inst_names: *const *const std::ffi::c_char,
length: usize,
) -> std::ffi::c_int {
if in_config_server_callback() {
set_error_msg("cannot delete network instances from config server callback");
return -1;
}
wait_for_config_server_delivery();
let _remote_mutation_guard = lock_remote_instance_mutation();
let _mutation_guard = match INSTANCE_MUTATION_LOCK.lock() {
Ok(guard) => guard,
Err(err) => {
set_error_msg(&format!("failed to lock instance mutation: {}", err));
return -1;
}
};
if length == 0 {
return 0;
}
let Some(inst_names) = (unsafe { parse_instance_names(inst_names, length) }) else {
return -1;
};
let removed_ids = inst_names
.iter()
.filter_map(|name| INSTANCE_NAME_ID_MAP.get(name).map(|id| *id.value()))
.collect::<Vec<_>>();
if let Err(e) = INSTANCE_MANAGER.delete_network_instance(removed_ids.clone()) {
set_error_msg(&format!("failed to delete instances: {}", e));
return -1;
}
remove_config_server_tracked_instance_ids(&removed_ids);
remove_data_plane_handles_by_instance_ids(&removed_ids);
for name in inst_names {
INSTANCE_NAME_ID_MAP.remove(&name);
}
0
}
/// # Safety
/// Collect the network infos
pub(crate) unsafe fn collect_network_infos(
infos: *mut KeyValuePair,
max_length: usize,
) -> std::ffi::c_int {
if in_config_server_callback() {
set_error_msg("cannot collect network infos from config server callback");
return -1;
}
if max_length == 0 {
return 0;
}
let infos = unsafe {
assert!(!infos.is_null());
std::slice::from_raw_parts_mut(infos, max_length)
};
let collected_infos = match INSTANCE_MANAGER.collect_network_infos_sync() {
Ok(infos) => infos,
Err(e) => {
set_error_msg(&format!("failed to collect network infos: {}", e));
return -1;
}
};
let mut index = 0;
for (instance_id, value) in collected_infos.iter() {
if index >= max_length {
break;
}
let Some(key) = INSTANCE_MANAGER.get_instance_name(instance_id) else {
continue;
};
// convert value to json string
let value = match serde_json::to_string(&value) {
Ok(value) => value,
Err(e) => {
set_error_msg(&format!("failed to serialize instance info: {}", e));
return -1;
}
};
infos[index] = KeyValuePair {
key: std::ffi::CString::new(key).unwrap().into_raw(),
value: std::ffi::CString::new(value).unwrap().into_raw(),
};
index += 1;
}
index as std::ffi::c_int
}
/// # Safety
/// List the instance names and IDs known by the FFI instance manager.
pub(crate) unsafe fn list_instance(infos: *mut KeyValuePair, max_length: usize) -> std::ffi::c_int {
if in_config_server_callback() {
set_error_msg("cannot list instances from config server callback");
return -1;
}
if max_length == 0 {
return 0;
}
if infos.is_null() {
set_error_msg("infos is null");
return -1;
}
let infos = unsafe { std::slice::from_raw_parts_mut(infos, max_length) };
let mut instances = INSTANCE_MANAGER
.list_network_instance_ids()
.into_iter()
.filter_map(|id| {
INSTANCE_MANAGER
.get_instance_name(&id)
.map(|name| (name, id))
})
.collect::<Vec<_>>();
instances.sort_by(|(left_name, left_id), (right_name, right_id)| {
left_name
.cmp(right_name)
.then_with(|| left_id.to_string().cmp(&right_id.to_string()))
});
let encoded_instances = match instances
.into_iter()
.take(max_length)
.map(|(name, id)| {
let key = CString::new(name)
.map_err(|err| format!("failed to encode instance name: {}", err))?;
let value = CString::new(id.to_string())
.map_err(|err| format!("failed to encode instance id: {}", err))?;
Ok((key, value))
})
.collect::<Result<Vec<_>, String>>()
{
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let count = encoded_instances.len();
for (index, (key, value)) in encoded_instances.into_iter().enumerate() {
infos[index] = KeyValuePair {
key: key.into_raw(),
value: value.into_raw(),
};
}
count as std::ffi::c_int
}
@@ -0,0 +1,100 @@
use std::ffi::{CString, c_char, c_int};
use crate::{
config_server::in_config_server_callback,
error::set_error_msg,
state::{ASYNC_RUNTIME, INSTANCE_MANAGER},
strings::{c_str_to_string, optional_c_str_to_string},
};
/// # Safety
/// See `crate::call_json_rpc`.
pub(crate) unsafe fn call_json_rpc(
service_name: *const c_char,
method_name: *const c_char,
domain_name: *const c_char,
payload_json: *const c_char,
out_response_json: *mut *const c_char,
) -> c_int {
if out_response_json.is_null() {
set_error_msg("out_response_json is null");
return -1;
}
unsafe {
*out_response_json = std::ptr::null();
}
if in_config_server_callback() {
set_error_msg("cannot call JSON RPC from config server callback");
return -1;
}
let service_name = match unsafe { c_str_to_string(service_name, "service_name") } {
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let method_name = match unsafe { c_str_to_string(method_name, "method_name") } {
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let domain_name = match unsafe { optional_c_str_to_string(domain_name, "domain_name") } {
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let payload_json = match unsafe { c_str_to_string(payload_json, "payload_json") } {
Ok(value) => value,
Err(err) => {
set_error_msg(&err);
return -1;
}
};
let payload = match serde_json::from_str::<serde_json::Value>(&payload_json) {
Ok(value) => value,
Err(err) => {
set_error_msg(&format!("failed to parse payload_json: {}", err));
return -1;
}
};
let response = match ASYNC_RUNTIME.block_on(easytier::rpc_service::call_json_rpc(
&INSTANCE_MANAGER,
&service_name,
&method_name,
domain_name.as_deref(),
payload,
)) {
Ok(value) => value,
Err(err) => {
set_error_msg(&format!("RPC Error: {}", err));
return -1;
}
};
let response_json = match serde_json::to_string(&response) {
Ok(value) => value,
Err(err) => {
set_error_msg(&format!("failed to serialize RPC response: {}", err));
return -1;
}
};
let response_json = match CString::new(response_json) {
Ok(value) => value,
Err(err) => {
set_error_msg(&format!("failed to allocate RPC response: {}", err));
return -1;
}
};
unsafe {
*out_response_json = response_json.into_raw();
}
0
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,54 @@
use std::sync::{Arc, Mutex};
use dashmap::DashMap;
use easytier::instance_manager::NetworkInstanceManager;
use tokio::runtime::{Builder, Runtime};
use uuid::Uuid;
pub(crate) static INSTANCE_NAME_ID_MAP: once_cell::sync::Lazy<DashMap<String, Uuid>> =
once_cell::sync::Lazy::new(DashMap::new);
pub(crate) static INSTANCE_MANAGER: once_cell::sync::Lazy<Arc<NetworkInstanceManager>> =
once_cell::sync::Lazy::new(|| Arc::new(NetworkInstanceManager::new()));
pub(crate) 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-ffi")
});
pub(crate) static INSTANCE_MUTATION_LOCK: once_cell::sync::Lazy<Mutex<()>> =
once_cell::sync::Lazy::new(|| Mutex::new(()));
pub(crate) fn remove_instance_name_ids(ids: &[Uuid]) {
if ids.is_empty() {
return;
}
INSTANCE_NAME_ID_MAP.retain(|_, instance_id| !ids.contains(instance_id));
}
pub(crate) fn lock_remote_instance_mutation() -> tokio::sync::OwnedMutexGuard<()> {
INSTANCE_MANAGER
.remote_mutation_lock()
.blocking_lock_owned()
}
pub(crate) fn instance_name_exists(inst_name: &str) -> bool {
find_instance_id_by_name(inst_name).is_some()
}
pub(crate) fn find_instance_id_by_name(inst_name: &str) -> Option<Uuid> {
INSTANCE_NAME_ID_MAP
.get(inst_name)
.map(|id| *id)
.or_else(|| {
INSTANCE_MANAGER
.list_network_instance_ids()
.into_iter()
.find(|id| {
INSTANCE_MANAGER
.get_instance_name(id)
.is_some_and(|name| name == inst_name)
})
})
}
@@ -0,0 +1,23 @@
use std::ffi::{CStr, c_char};
pub(crate) unsafe fn c_str_to_string(ptr: *const c_char, name: &str) -> Result<String, String> {
if ptr.is_null() {
return Err(format!("{} is null", name));
}
unsafe { CStr::from_ptr(ptr) }
.to_str()
.map(|value| value.to_string())
.map_err(|err| format!("{} is not valid UTF-8: {}", name, err))
}
pub(crate) unsafe fn optional_c_str_to_string(
ptr: *const c_char,
name: &str,
) -> Result<Option<String>, String> {
if ptr.is_null() {
return Ok(None);
}
unsafe { c_str_to_string(ptr, name) }.map(Some)
}
+766
View File
@@ -0,0 +1,766 @@
use crate::{
config_server::{
ConfigServerCallbackScope, ManagedConfigServerClientHooks, set_active_for_test,
},
state::{
INSTANCE_MANAGER, INSTANCE_NAME_ID_MAP, find_instance_id_by_name,
lock_remote_instance_mutation, remove_instance_name_ids,
},
*,
};
use easytier::{
common::config::{ConfigFileControl, ConfigLoader as _, TomlConfigLoader},
web_client::WebClientHooks,
};
use serde_json::Value;
use std::{
collections::HashSet,
ffi::{CStr, CString, c_char, c_void},
sync::{Mutex, mpsc},
time::Duration,
};
use uuid::Uuid;
#[test]
fn test_parse_config() {
let cfg_str = r#"
inst_name = "test"
network = "test_network"
"#;
let cstr = std::ffi::CString::new(cfg_str).unwrap();
unsafe {
assert_eq!(parse_config(cstr.as_ptr()), 0);
}
}
#[test]
fn test_run_network_instance() {
let cfg_str = r#"
inst_name = "test"
network = "test_network"
"#;
let cstr = std::ffi::CString::new(cfg_str).unwrap();
unsafe {
assert_eq!(run_network_instance(cstr.as_ptr()), 0);
}
}
#[test]
fn get_error_msg_returns_config_server_callback_error() {
let hooks = ManagedConfigServerClientHooks::new(None, std::ptr::null_mut());
let callback_error = format!("callback delivery failed {}", Uuid::new_v4());
crate::config_server::clear_last_callback_error();
hooks.note_callback_error(callback_error.clone());
unsafe {
let mut error_ptr: *const c_char = std::ptr::null();
get_error_msg(&mut error_ptr);
assert!(!error_ptr.is_null());
let error_msg = CStr::from_ptr(error_ptr).to_string_lossy().into_owned();
free_string(error_ptr);
assert!(error_msg.contains(&callback_error));
}
crate::config_server::clear_last_callback_error();
}
unsafe extern "C" fn record_config_server_event(event_json: *const c_char, user_data: *mut c_void) {
let events = unsafe { &*(user_data as *const Mutex<Vec<String>>) };
events.lock().unwrap().push(
unsafe { CStr::from_ptr(event_json) }
.to_string_lossy()
.into_owned(),
);
}
fn take_last_error() -> Option<String> {
unsafe {
let mut error_ptr: *const c_char = std::ptr::null();
get_error_msg(&mut error_ptr);
if error_ptr.is_null() {
None
} else {
let error = CStr::from_ptr(error_ptr).to_string_lossy().into_owned();
free_string(error_ptr);
Some(error)
}
}
}
fn free_key_value_pairs(infos: &[KeyValuePair]) {
for info in infos {
free_string(info.key);
free_string(info.value);
}
}
#[test]
fn list_instance_returns_instance_names_and_ids() {
let instance_id = Uuid::new_v4();
let instance_name = format!("list-instance-{}", instance_id);
let cfg = TomlConfigLoader::default();
cfg.set_id(instance_id);
cfg.set_inst_name(instance_name.clone());
INSTANCE_MANAGER
.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG)
.unwrap();
INSTANCE_NAME_ID_MAP.insert(instance_name.clone(), instance_id);
let mut infos = vec![
KeyValuePair {
key: std::ptr::null(),
value: std::ptr::null(),
};
16
];
let count = unsafe { list_instance(infos.as_mut_ptr(), infos.len()) };
assert!(count > 0);
let mut found = false;
for info in infos.iter().take(count as usize) {
let key = unsafe { CStr::from_ptr(info.key) }.to_string_lossy();
let value = unsafe { CStr::from_ptr(info.value) }.to_string_lossy();
if key == instance_name {
assert_eq!(value, instance_id.to_string());
found = true;
}
}
free_key_value_pairs(&infos[..count as usize]);
INSTANCE_MANAGER
.delete_network_instance(vec![instance_id])
.unwrap();
remove_instance_name_ids(&[instance_id]);
assert!(found);
}
#[test]
fn list_instance_allows_zero_length() {
assert_eq!(unsafe { list_instance(std::ptr::null_mut(), 0) }, 0);
}
#[test]
fn list_instance_rejects_null_output_pointer() {
assert_eq!(unsafe { list_instance(std::ptr::null_mut(), 1) }, -1);
assert!(take_last_error().unwrap().contains("infos is null"));
}
#[test]
fn call_json_rpc_returns_logger_response() {
let service = CString::new("api.logger.LoggerRpcService").unwrap();
let method = CString::new("get_logger_config").unwrap();
let payload = CString::new("{}").unwrap();
let mut response_ptr: *const c_char = std::ptr::null();
assert_eq!(
unsafe {
call_json_rpc(
service.as_ptr(),
method.as_ptr(),
std::ptr::null(),
payload.as_ptr(),
&mut response_ptr,
)
},
0
);
assert!(!response_ptr.is_null());
let response = unsafe { CStr::from_ptr(response_ptr) }
.to_string_lossy()
.into_owned();
free_string(response_ptr);
let response: Value = serde_json::from_str(&response).unwrap();
assert!(response.get("level").is_some());
}
#[test]
fn call_json_rpc_rejects_instance_management_service() {
let service = CString::new("api.manage.WebClientService").unwrap();
let method = CString::new("list_network_instance").unwrap();
let payload = CString::new("{}").unwrap();
let mut response_ptr: *const c_char = std::ptr::null();
assert_eq!(
unsafe {
call_json_rpc(
service.as_ptr(),
method.as_ptr(),
std::ptr::null(),
payload.as_ptr(),
&mut response_ptr,
)
},
-1
);
assert!(response_ptr.is_null());
assert!(take_last_error().unwrap().contains("not exposed"));
}
#[test]
fn call_json_rpc_rejects_malformed_payload_json() {
let service = CString::new("api.logger.LoggerRpcService").unwrap();
let method = CString::new("get_logger_config").unwrap();
let payload = CString::new("{").unwrap();
let mut response_ptr: *const c_char = std::ptr::null();
assert_eq!(
unsafe {
call_json_rpc(
service.as_ptr(),
method.as_ptr(),
std::ptr::null(),
payload.as_ptr(),
&mut response_ptr,
)
},
-1
);
assert!(response_ptr.is_null());
assert!(
take_last_error()
.unwrap()
.contains("failed to parse payload_json")
);
}
#[test]
fn call_json_rpc_rejects_null_output_pointer() {
let service = CString::new("api.logger.LoggerRpcService").unwrap();
let method = CString::new("get_logger_config").unwrap();
let payload = CString::new("{}").unwrap();
assert_eq!(
unsafe {
call_json_rpc(
service.as_ptr(),
method.as_ptr(),
std::ptr::null(),
payload.as_ptr(),
std::ptr::null_mut(),
)
},
-1
);
assert!(
take_last_error()
.unwrap()
.contains("out_response_json is null")
);
}
#[tokio::test]
async fn config_server_hooks_emit_run_event() {
let events: Mutex<Vec<String>> = Mutex::new(Vec::new());
let hooks = ManagedConfigServerClientHooks::new(
Some(record_config_server_event),
&events as *const _ as *mut c_void,
);
let instance_id = Uuid::new_v4();
let cfg = TomlConfigLoader::default();
cfg.set_id(instance_id);
let inst_name = format!("test-{}", instance_id);
cfg.set_inst_name(inst_name.clone());
hooks.pre_run_network_instance(&cfg).await.unwrap();
INSTANCE_MANAGER
.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG)
.unwrap();
hooks.post_run_network_instance(&instance_id).await.unwrap();
let duplicate_cfg = TomlConfigLoader::default();
duplicate_cfg.set_inst_name(inst_name);
duplicate_cfg.set_id(Uuid::new_v4());
assert!(
hooks
.pre_run_network_instance(&duplicate_cfg)
.await
.is_err()
);
assert_eq!(hooks.tracked_instance_ids(), vec![instance_id]);
let events = events.lock().unwrap();
assert_eq!(events.len(), 1);
let event: Value = serde_json::from_str(&events[0]).unwrap();
assert_eq!(event["event"], "run_network_instance");
assert_eq!(event["success"], true);
assert_eq!(event["instance_id"], instance_id.to_string());
assert!(event["error"].is_null());
INSTANCE_MANAGER
.delete_network_instance(vec![instance_id])
.unwrap();
remove_instance_name_ids(&[instance_id]);
}
#[tokio::test]
async fn config_server_hooks_emit_delete_events_for_tracked_instances() {
let events: Mutex<Vec<String>> = Mutex::new(Vec::new());
let hooks = ManagedConfigServerClientHooks::new(
Some(record_config_server_event),
&events as *const _ as *mut c_void,
);
let instance_id_1 = Uuid::new_v4();
let instance_id_2 = Uuid::new_v4();
let unknown_instance_id = Uuid::new_v4();
for id in [instance_id_1, instance_id_2] {
let cfg = TomlConfigLoader::default();
cfg.set_id(id);
cfg.set_inst_name(format!("test-{}", id));
hooks.pre_run_network_instance(&cfg).await.unwrap();
INSTANCE_MANAGER
.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG)
.unwrap();
}
hooks
.post_run_network_instance(&instance_id_1)
.await
.unwrap();
hooks
.post_run_network_instance(&instance_id_2)
.await
.unwrap();
events.lock().unwrap().clear();
hooks
.post_remove_network_instances(&[instance_id_1, unknown_instance_id, instance_id_2])
.await
.unwrap();
assert!(hooks.tracked_instance_ids().is_empty());
let events = events.lock().unwrap();
assert_eq!(events.len(), 2);
let event_ids = events
.iter()
.map(|event| {
let event: Value = serde_json::from_str(event).unwrap();
assert_eq!(event["event"], "delete_network_instance");
assert_eq!(event["success"], true);
assert!(event["error"].is_null());
event["instance_id"].as_str().unwrap().to_string()
})
.collect::<HashSet<_>>();
assert_eq!(
event_ids,
HashSet::from([instance_id_1.to_string(), instance_id_2.to_string()])
);
INSTANCE_MANAGER
.delete_network_instance(vec![instance_id_1, instance_id_2])
.unwrap();
remove_instance_name_ids(&[instance_id_1, instance_id_2]);
}
#[tokio::test]
async fn config_server_hooks_remove_untracked_name_mapping_without_event() {
let events: Mutex<Vec<String>> = Mutex::new(Vec::new());
let hooks = ManagedConfigServerClientHooks::new(
Some(record_config_server_event),
&events as *const _ as *mut c_void,
);
let local_id = Uuid::new_v4();
let inst_name = format!("local-{}", local_id);
INSTANCE_NAME_ID_MAP.insert(inst_name.clone(), local_id);
hooks
.post_remove_network_instances(&[local_id])
.await
.unwrap();
assert!(INSTANCE_NAME_ID_MAP.get(&inst_name).is_none());
assert!(events.lock().unwrap().is_empty());
}
#[tokio::test]
async fn config_server_hooks_reject_duplicate_instance_name() {
let hooks = ManagedConfigServerClientHooks::new(None, std::ptr::null_mut());
let inst_name = format!("test-{}", Uuid::new_v4());
let existing_id = Uuid::new_v4();
let new_id = Uuid::new_v4();
INSTANCE_NAME_ID_MAP.insert(inst_name.clone(), existing_id);
let cfg = TomlConfigLoader::default();
cfg.set_inst_name(inst_name.clone());
cfg.set_id(new_id);
assert!(hooks.pre_run_network_instance(&cfg).await.is_err());
assert_eq!(*INSTANCE_NAME_ID_MAP.get(&inst_name).unwrap(), existing_id);
INSTANCE_NAME_ID_MAP.remove(&inst_name);
}
#[tokio::test]
async fn config_server_hooks_remove_overwritten_id_before_duplicate_name_error() {
let events: Mutex<Vec<String>> = Mutex::new(Vec::new());
let hooks = ManagedConfigServerClientHooks::new(
Some(record_config_server_event),
&events as *const _ as *mut c_void,
);
let old_name = format!("old-{}", Uuid::new_v4());
let duplicate_name = format!("duplicate-{}", Uuid::new_v4());
let overwritten_id = Uuid::new_v4();
let duplicate_id = Uuid::new_v4();
hooks.instance_ids.lock().unwrap().insert(overwritten_id);
INSTANCE_NAME_ID_MAP.insert(old_name.clone(), overwritten_id);
INSTANCE_NAME_ID_MAP.insert(duplicate_name.clone(), duplicate_id);
hooks
.post_remove_network_instances(&[overwritten_id])
.await
.unwrap();
let cfg = TomlConfigLoader::default();
cfg.set_inst_name(duplicate_name.clone());
cfg.set_id(overwritten_id);
assert!(hooks.pre_run_network_instance(&cfg).await.is_err());
assert!(hooks.tracked_instance_ids().is_empty());
assert!(INSTANCE_NAME_ID_MAP.get(&old_name).is_none());
assert_eq!(
*INSTANCE_NAME_ID_MAP.get(&duplicate_name).unwrap(),
duplicate_id
);
assert_eq!(events.lock().unwrap().len(), 1);
INSTANCE_NAME_ID_MAP.remove(&duplicate_name);
}
#[tokio::test]
async fn config_server_hooks_remove_tracked_state_before_overwrite_retry() {
let hooks = ManagedConfigServerClientHooks::new(None, std::ptr::null_mut());
let inst_name = format!("test-{}", Uuid::new_v4());
let instance_id = Uuid::new_v4();
hooks.instance_ids.lock().unwrap().insert(instance_id);
INSTANCE_NAME_ID_MAP.insert(inst_name.clone(), instance_id);
let cfg = TomlConfigLoader::default();
cfg.set_inst_name(inst_name.clone());
cfg.set_id(instance_id);
hooks
.post_remove_network_instances(&[instance_id])
.await
.unwrap();
hooks.pre_run_network_instance(&cfg).await.unwrap();
assert!(hooks.tracked_instance_ids().is_empty());
assert!(INSTANCE_NAME_ID_MAP.get(&inst_name).is_none());
}
#[tokio::test]
async fn config_server_hooks_reject_post_run_after_external_delete() {
let hooks = ManagedConfigServerClientHooks::new(None, std::ptr::null_mut());
let instance_id = Uuid::new_v4();
let cfg = TomlConfigLoader::default();
cfg.set_id(instance_id);
cfg.set_inst_name(format!("test-{}", instance_id));
hooks.pre_run_network_instance(&cfg).await.unwrap();
INSTANCE_MANAGER
.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG)
.unwrap();
INSTANCE_MANAGER
.delete_network_instance(vec![instance_id])
.unwrap();
assert!(hooks.post_run_network_instance(&instance_id).await.is_err());
}
#[test]
fn find_instance_id_by_name_resolves_uncommitted_manager_instance_name() {
let instance_id = Uuid::new_v4();
let inst_name = format!("test-{}", instance_id);
let cfg = TomlConfigLoader::default();
cfg.set_id(instance_id);
cfg.set_inst_name(inst_name.clone());
INSTANCE_MANAGER
.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG)
.unwrap();
assert_eq!(find_instance_id_by_name(&inst_name), Some(instance_id));
INSTANCE_MANAGER
.delete_network_instance(vec![instance_id])
.unwrap();
remove_instance_name_ids(&[instance_id]);
}
#[test]
fn delete_network_instance_removes_only_named_instances() {
let keep_id = Uuid::new_v4();
let delete_id = Uuid::new_v4();
let keep_name = format!("keep-{}", keep_id);
let delete_name = format!("delete-{}", delete_id);
for (id, name) in [
(keep_id, keep_name.clone()),
(delete_id, delete_name.clone()),
] {
let cfg = TomlConfigLoader::default();
cfg.set_id(id);
cfg.set_inst_name(name.clone());
INSTANCE_MANAGER
.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG)
.unwrap();
INSTANCE_NAME_ID_MAP.insert(name, id);
}
let delete_name = CString::new(delete_name.clone()).unwrap();
let inst_names = [delete_name.as_ptr()];
assert_eq!(
unsafe { delete_network_instance(inst_names.as_ptr(), inst_names.len()) },
0
);
assert_eq!(find_instance_id_by_name(&keep_name), Some(keep_id));
assert!(find_instance_id_by_name(delete_name.to_str().unwrap()).is_none());
INSTANCE_MANAGER
.delete_network_instance(vec![keep_id])
.unwrap();
remove_instance_name_ids(&[keep_id]);
}
#[test]
fn retain_and_delete_network_instance_reject_invalid_name_pointers() {
assert_eq!(unsafe { retain_network_instance(std::ptr::null(), 1) }, -1);
assert_eq!(unsafe { delete_network_instance(std::ptr::null(), 1) }, -1);
let inst_names = [std::ptr::null()];
assert_eq!(
unsafe { retain_network_instance(inst_names.as_ptr(), inst_names.len()) },
-1
);
assert_eq!(
unsafe { delete_network_instance(inst_names.as_ptr(), inst_names.len()) },
-1
);
}
#[test]
fn ffi_remote_mutation_lock_uses_manager_lock() {
let manager_guard = INSTANCE_MANAGER
.remote_mutation_lock()
.blocking_lock_owned();
let (done_tx, done_rx) = mpsc::channel();
let waiter = std::thread::spawn(move || {
let _ffi_guard = lock_remote_instance_mutation();
done_tx.send(()).unwrap();
});
assert!(done_rx.recv_timeout(Duration::from_millis(100)).is_err());
drop(manager_guard);
done_rx.recv_timeout(Duration::from_secs(5)).unwrap();
waiter.join().unwrap();
}
#[tokio::test]
async fn config_server_hooks_suppress_late_run_events_while_stopping() {
let events: Mutex<Vec<String>> = Mutex::new(Vec::new());
let hooks = ManagedConfigServerClientHooks::new(
Some(record_config_server_event),
&events as *const _ as *mut c_void,
);
hooks.start_stopping();
hooks
.post_run_network_instance(&Uuid::new_v4())
.await
.unwrap();
assert!(hooks.tracked_instance_ids().is_empty());
assert!(events.lock().unwrap().is_empty());
}
#[test]
fn config_server_callback_context_rejects_nested_blocking_ffi_calls() {
let _callback_scope = ConfigServerCallbackScope::enter();
assert_eq!(is_config_server_client_connected(), 0);
let service = CString::new("api.logger.LoggerRpcService").unwrap();
let method = CString::new("get_logger_config").unwrap();
let payload = CString::new("{}").unwrap();
let mut response_ptr: *const c_char = std::ptr::null();
assert_eq!(
unsafe {
call_json_rpc(
service.as_ptr(),
method.as_ptr(),
std::ptr::null(),
payload.as_ptr(),
&mut response_ptr,
)
},
-1
);
assert!(response_ptr.is_null());
assert_eq!(
unsafe { collect_network_infos(std::ptr::null_mut(), 0) },
-1
);
assert_eq!(unsafe { list_instance(std::ptr::null_mut(), 0) }, -1);
let cfg = CString::new("inst_name = \"callback-test\"\nlisteners = []").unwrap();
assert_eq!(unsafe { run_network_instance(cfg.as_ptr()) }, -1);
assert_eq!(unsafe { retain_network_instance(std::ptr::null(), 0) }, -1);
assert_eq!(unsafe { delete_network_instance(std::ptr::null(), 0) }, -1);
let url = CString::new("ring://test/token").unwrap();
let machine_id = CString::new("test-machine").unwrap();
assert_eq!(
unsafe {
start_config_server_client(
url.as_ptr(),
std::ptr::null(),
machine_id.as_ptr(),
false,
None,
std::ptr::null_mut(),
)
},
-1
);
assert_eq!(stop_config_server_client(), -1);
#[cfg(feature = "ffi-dataplane")]
{
assert_eq!(
unsafe {
data_plane_tcp_connect(
std::ptr::null(),
std::ptr::null(),
0,
0,
std::ptr::null_mut(),
std::ptr::null_mut(),
)
},
0
);
assert_eq!(
unsafe {
data_plane_tcp_bind(
std::ptr::null(),
0,
0,
std::ptr::null_mut(),
std::ptr::null_mut(),
)
},
0
);
assert_eq!(
unsafe {
data_plane_tcp_accept(
0,
0,
std::ptr::null_mut(),
std::ptr::null_mut(),
std::ptr::null_mut(),
std::ptr::null_mut(),
)
},
0
);
assert_eq!(
unsafe { data_plane_tcp_read(0, std::ptr::null_mut(), 0, 0) },
-1
);
assert_eq!(
unsafe { data_plane_tcp_write(0, std::ptr::null(), 0, 0) },
-1
);
assert_eq!(data_plane_tcp_close(0), -1);
assert_eq!(data_plane_tcp_listener_close(0), -1);
assert_eq!(
unsafe {
data_plane_udp_bind(
std::ptr::null(),
0,
0,
std::ptr::null_mut(),
std::ptr::null_mut(),
)
},
0
);
assert_eq!(
unsafe { data_plane_udp_send_to(0, std::ptr::null(), 0, std::ptr::null(), 0, 0) },
-1
);
assert_eq!(
unsafe {
data_plane_udp_recv_from(
0,
std::ptr::null_mut(),
0,
std::ptr::null_mut(),
std::ptr::null_mut(),
0,
)
},
-1
);
assert_eq!(data_plane_udp_close(0), -1);
assert_eq!(data_plane_async_op_status(0), -2);
assert_eq!(data_plane_async_op_wait(0, 0), -2);
assert_eq!(data_plane_async_op_cancel(0), -2);
assert_eq!(data_plane_async_op_free(0), -2);
data_plane_free_bytes(std::ptr::null(), 0);
assert_eq!(
unsafe { data_plane_tcp_connect_start(std::ptr::null(), std::ptr::null(), 0, 0) },
0
);
assert_eq!(
unsafe { data_plane_tcp_bind_start(std::ptr::null(), 0, 0) },
0
);
assert_eq!(unsafe { data_plane_tcp_accept_start(0, 0) }, 0);
assert_eq!(unsafe { data_plane_tcp_read_start(0, 0, 0) }, 0);
assert_eq!(
unsafe { data_plane_tcp_write_start(0, std::ptr::null(), 0, 0) },
0
);
assert_eq!(
unsafe { data_plane_udp_bind_start(std::ptr::null(), 0, 0) },
0
);
assert_eq!(
unsafe { data_plane_udp_send_to_start(0, std::ptr::null(), 0, std::ptr::null(), 0, 0) },
0
);
assert_eq!(unsafe { data_plane_udp_recv_from_start(0, 0, 0) }, 0);
}
}
#[cfg(feature = "ffi-dataplane")]
#[test]
fn active_config_server_rejects_data_plane() {
set_active_for_test(true);
assert_eq!(
unsafe {
data_plane_tcp_connect(
std::ptr::null(),
std::ptr::null(),
0,
0,
std::ptr::null_mut(),
std::ptr::null_mut(),
)
},
0
);
assert_eq!(
unsafe { data_plane_tcp_read(0, std::ptr::null_mut(), 0, 0) },
-1
);
assert_eq!(
unsafe { data_plane_tcp_connect_start(std::ptr::null(), std::ptr::null(), 0, 0) },
0
);
assert_eq!(unsafe { data_plane_tcp_read_start(0, 0, 0) }, 0);
set_active_for_test(false);
}
#[cfg(feature = "ffi-dataplane")]
#[test]
fn async_op_invalid_handle_helpers_are_stable() {
assert_eq!(data_plane_async_op_status(u64::MAX), -2);
assert_eq!(data_plane_async_op_wait(u64::MAX, 1), -2);
assert_eq!(data_plane_async_op_cancel(u64::MAX), -2);
assert_eq!(data_plane_async_op_free(u64::MAX), -2);
data_plane_free_bytes(std::ptr::null(), 0);
}
@@ -0,0 +1,10 @@
use std::ffi::{c_char, c_void};
#[repr(C)]
#[derive(Clone, Copy)]
pub struct KeyValuePair {
pub key: *const c_char,
pub value: *const c_char,
}
pub type ConfigServerEventCallback = Option<unsafe extern "C" fn(*const c_char, *mut c_void)>;
+113 -65
View File
@@ -1237,18 +1237,18 @@ dependencies = [
"ordered_hash_map",
"parking_lot",
"paste",
"pbjson",
"pbjson-build",
"percent-encoding",
"petgraph 0.8.2",
"petgraph",
"pin-project-lite",
"pnet",
"prefix-trie",
"proc-macro2",
"prost",
"prost 0.14.3",
"prost-build",
"prost-reflect",
"prost-reflect 0.16.4",
"prost-reflect-build",
"prost-wkt",
"prost-wkt-build",
"prost-wkt-types",
"quinn",
"quinn-plaintext",
@@ -1318,9 +1318,8 @@ dependencies = [
"napi-build-ohos",
"napi-derive-ohos",
"napi-ohos",
"ohos-hilog-binding",
"once_cell",
"prost-reflect",
"prost-reflect 0.14.7",
"rusqlite",
"serde",
"serde_json",
@@ -2595,7 +2594,7 @@ dependencies = [
[[package]]
name = "kcp-sys"
version = "0.1.0"
source = "git+https://github.com/EasyTier/kcp-sys?rev=94964794caaed5d388463137da59b97499619e5f#94964794caaed5d388463137da59b97499619e5f"
source = "git+https://github.com/EasyTier/kcp-sys?rev=d7427c22d764deb1860a7d37acc446ed5033464c#d7427c22d764deb1860a7d37acc446ed5033464c"
dependencies = [
"anyhow",
"auto_impl",
@@ -3165,22 +3164,6 @@ dependencies = [
"libc",
]
[[package]]
name = "ohos-hilog-binding"
version = "0.1.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "860d1e3c2c5e3217d819a16c815d2d4dcbc7610285d2612d08745a29c353a503"
dependencies = [
"libc",
"ohos-hilogs-sys",
]
[[package]]
name = "ohos-hilogs-sys"
version = "0.0.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ed07615005d0f8d7bcf901f89c8ff4870666a9bdb00382f588af383f40c160b7"
[[package]]
name = "once_cell"
version = "1.21.3"
@@ -3308,6 +3291,28 @@ version = "1.0.15"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a"
[[package]]
name = "pbjson"
version = "0.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e8edd1efdd8ab23ba9cb9ace3d9987a72663d5d7c9f74fa00b51d6213645cf6c"
dependencies = [
"base64 0.22.1",
"serde",
]
[[package]]
name = "pbjson-build"
version = "0.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2ed4d5c6ae95e08ac768883c8401cf0e8deb4e6e1d6a4e1fd3d2ec4f0ec63200"
dependencies = [
"heck 0.5.0",
"itertools 0.14.0",
"prost 0.14.3",
"prost-types 0.14.3",
]
[[package]]
name = "pbkdf2"
version = "0.12.2"
@@ -3334,16 +3339,6 @@ version = "2.3.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220"
[[package]]
name = "petgraph"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3672b37090dbd86368a4145bc067582552b29c27377cad4e0a306c97f9bd7772"
dependencies = [
"fixedbitset",
"indexmap",
]
[[package]]
name = "petgraph"
version = "0.8.2"
@@ -3631,24 +3626,33 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2796faa41db3ec313a31f7624d9286acf277b52de526150b7e69f3debf891ee5"
dependencies = [
"bytes",
"prost-derive",
"prost-derive 0.13.5",
]
[[package]]
name = "prost"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d2ea70524a2f82d518bce41317d0fae74151505651af45faf1ffbd6fd33f0568"
dependencies = [
"bytes",
"prost-derive 0.14.3",
]
[[package]]
name = "prost-build"
version = "0.13.5"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "be769465445e8c1474e9c5dac2018218498557af32d9ed057325ec9a41ae81bf"
checksum = "343d3bd7056eda839b03204e68deff7d1b13aba7af2b2fd16890697274262ee7"
dependencies = [
"heck 0.5.0",
"itertools 0.14.0",
"log",
"multimap",
"once_cell",
"petgraph 0.7.1",
"petgraph",
"prettyplease",
"prost",
"prost-types",
"prost 0.14.3",
"prost-types 0.14.3",
"regex",
"syn 2.0.106",
"tempfile",
@@ -3667,6 +3671,19 @@ dependencies = [
"syn 2.0.106",
]
[[package]]
name = "prost-derive"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "27c6023962132f4b30eb4c172c91ce92d933da334c59c23cddee82358ddafb0b"
dependencies = [
"anyhow",
"itertools 0.14.0",
"proc-macro2",
"quote",
"syn 2.0.106",
]
[[package]]
name = "prost-reflect"
version = "0.14.7"
@@ -3674,19 +3691,30 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b5edd582b62f5cde844716e66d92565d7faf7ab1445c8cebce6e00fba83ddb2"
dependencies = [
"once_cell",
"prost",
"prost-reflect-derive",
"prost-types",
"prost 0.13.5",
"prost-reflect-derive 0.14.0",
"prost-types 0.13.5",
]
[[package]]
name = "prost-reflect"
version = "0.16.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "590aa145fee8f7a26b5a6055365e7c5e89a5c1caae9869de76ec0ee73181a2f9"
dependencies = [
"prost 0.14.3",
"prost-reflect-derive 0.16.0",
"prost-types 0.14.3",
]
[[package]]
name = "prost-reflect-build"
version = "0.14.0"
version = "0.16.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "50e2537231d94dd2778920c2ada37dd9eb1ac0325bb3ee3ee651bd44c1134123"
checksum = "8214ae2c30bbac390db0134d08300e770ef89b6d4e5abf855e8d300eded87e28"
dependencies = [
"prost-build",
"prost-reflect",
"prost-reflect 0.16.4",
]
[[package]]
@@ -3700,24 +3728,44 @@ dependencies = [
"syn 2.0.106",
]
[[package]]
name = "prost-reflect-derive"
version = "0.16.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7b6d90e29fa6c0d13c2c19ba5e4b3fb0efbf5975d27bcf4e260b7b15455bcabe"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.106",
]
[[package]]
name = "prost-types"
version = "0.13.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "52c2c1bf36ddb1a1c396b3601a3cec27c2462e45f07c386894ec3ccf5332bd16"
dependencies = [
"prost",
"prost 0.13.5",
]
[[package]]
name = "prost-types"
version = "0.14.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8991c4cbdb8bc5b11f0b074ffe286c30e523de90fee5ba8132f1399f23cb3dd7"
dependencies = [
"prost 0.14.3",
]
[[package]]
name = "prost-wkt"
version = "0.6.1"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "497e1e938f0c09ef9cabe1d49437b4016e03e8f82fbbe5d1c62a9b61b9decae1"
checksum = "cd3de5e9c9e84fcb5efa204b8e283d23e615a8bc8c777bf1d6622bb01dc61445"
dependencies = [
"chrono",
"inventory",
"prost",
"prost 0.14.3",
"serde",
"serde_derive",
"serde_json",
@@ -3726,27 +3774,27 @@ dependencies = [
[[package]]
name = "prost-wkt-build"
version = "0.6.1"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "07b8bf115b70a7aa5af1fd5d6e9418492e9ccb6e4785e858c938e28d132a884b"
checksum = "fe500dc80e757a75e1e8fb7290e448d62dfba3105ece1d058579cb00b58151cd"
dependencies = [
"heck 0.5.0",
"prost",
"prost 0.14.3",
"prost-build",
"prost-types",
"prost-types 0.14.3",
"quote",
]
[[package]]
name = "prost-wkt-types"
version = "0.6.1"
version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c8cdde6df0a98311c839392ca2f2f0bcecd545f86a62b4e3c6a49c336e970fe5"
checksum = "13807eaa7e15833d06e899008371926201cdcd11d74b6d490f49130cdb3f415e"
dependencies = [
"chrono",
"prost",
"prost 0.14.3",
"prost-build",
"prost-types",
"prost-types 0.14.3",
"prost-wkt",
"prost-wkt-build",
"regex",
@@ -3835,9 +3883,9 @@ dependencies = [
[[package]]
name = "quote"
version = "1.0.40"
version = "1.0.45"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1885c039570dc00dcb4ff087a89e185fd56bae234ddc7f056a945bf36467248d"
checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924"
dependencies = [
"proc-macro2",
]
@@ -3974,9 +4022,9 @@ dependencies = [
[[package]]
name = "regex"
version = "1.11.2"
version = "1.12.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "23d7fd106d8c02486a8d64e778353d1cffe08ce79ac2e82f540c86d0facf6912"
checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276"
dependencies = [
"aho-corasick",
"memchr",
@@ -3986,9 +4034,9 @@ dependencies = [
[[package]]
name = "regex-automata"
version = "0.4.10"
version = "0.4.14"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6b9458fa0bfeeac22b5ca447c63aaf45f28439a709ccd244698632f9aa6394d6"
checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f"
dependencies = [
"aho-corasick",
"memchr",
@@ -11,7 +11,6 @@ async-trait = "0.1"
base64 = "0.22"
flate2 = "1.1"
gethostname = "1.1"
ohos-hilog-binding = {version = "*", features = ["redirect"]}
easytier = { path = "../../easytier" }
napi-derive-ohos = "1.1"
napi-ohos = { version = "1.1", default-features = false, features = [
@@ -1,13 +1,49 @@
use crate::config::types::stored_config::{StoredConfigList, StoredConfigMeta};
use ohos_hilog_binding::{hilog_debug, hilog_error};
use crate::config::types::stored_config::{
SnapshotImportResult, StoredConfigList, StoredConfigMeta,
};
use once_cell::sync::Lazy;
use rusqlite::{Connection, OptionalExtension, params};
use std::collections::HashSet;
use std::ops::{Deref, DerefMut};
use std::path::{Path, PathBuf};
use std::sync::Mutex;
use std::sync::{Mutex, MutexGuard};
use std::time::{SystemTime, UNIX_EPOCH};
static CONFIG_DB_PATH: Mutex<Option<PathBuf>> = Mutex::new(None);
static CONFIG_DB_CONNECTION: Lazy<Mutex<Option<CachedConfigDb>>> = Lazy::new(|| Mutex::new(None));
const CONFIG_DB_FILE_NAME: &str = "easytier-config-store.db";
struct CachedConfigDb {
path: PathBuf,
conn: Connection,
}
pub(crate) struct ConfigDbGuard<'a> {
guard: MutexGuard<'a, Option<CachedConfigDb>>,
}
impl Deref for ConfigDbGuard<'_> {
type Target = Connection;
fn deref(&self) -> &Self::Target {
&self
.guard
.as_ref()
.expect("config db connection guard must contain a connection")
.conn
}
}
impl DerefMut for ConfigDbGuard<'_> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self
.guard
.as_mut()
.expect("config db connection guard must contain a connection")
.conn
}
}
#[derive(Debug, Clone)]
struct StoredConfigMetaRecord {
config_id: String,
@@ -18,6 +54,30 @@ struct StoredConfigMetaRecord {
temporary: bool,
}
type SnapshotFieldRow = (String, String, String, String);
fn snapshot_import_ok() -> SnapshotImportResult {
SnapshotImportResult {
ok: true,
error_code: String::new(),
error_message: String::new(),
snapshot_invalid: false,
}
}
fn snapshot_import_err(
error_code: &str,
error_message: impl Into<String>,
snapshot_invalid: bool,
) -> SnapshotImportResult {
SnapshotImportResult {
ok: false,
error_code: error_code.to_string(),
error_message: error_message.into(),
snapshot_invalid,
}
}
pub(crate) fn now_ts_string() -> String {
SystemTime::now()
.duration_since(UNIX_EPOCH)
@@ -53,31 +113,176 @@ fn init_schema(conn: &Connection) -> rusqlite::Result<()> {
);
CREATE INDEX IF NOT EXISTS idx_stored_config_fields_config_id
ON stored_config_fields(config_id);",
)
)?;
ensure_column(
conn,
"stored_configs",
"favorite",
"ALTER TABLE stored_configs ADD COLUMN favorite INTEGER NOT NULL DEFAULT 0;",
)?;
ensure_column(
conn,
"stored_configs",
"temporary",
"ALTER TABLE stored_configs ADD COLUMN temporary INTEGER NOT NULL DEFAULT 0;",
)?;
ensure_column(
conn,
"stored_config_fields",
"updated_at",
"ALTER TABLE stored_config_fields ADD COLUMN updated_at TEXT NOT NULL DEFAULT '0';",
)?;
if !validate_store_schema(conn)? {
return Err(rusqlite::Error::InvalidQuery);
}
conn.execute_batch("PRAGMA user_version = 1;")
}
pub(crate) fn open_db() -> Option<Connection> {
let path = db_file_path()?;
let conn = match Connection::open(&path) {
fn table_columns(conn: &Connection, table_name: &str) -> rusqlite::Result<HashSet<String>> {
let mut stmt = conn.prepare(&format!("PRAGMA table_info({})", table_name))?;
let rows = stmt.query_map([], |row| row.get::<_, String>(1))?;
let mut columns = HashSet::new();
for row in rows {
columns.insert(row?);
}
Ok(columns)
}
fn ensure_column(
conn: &Connection,
table_name: &str,
column_name: &str,
alter_sql: &str,
) -> rusqlite::Result<()> {
let columns = table_columns(conn, table_name)?;
if !columns.contains(column_name) {
conn.execute_batch(alter_sql)?;
}
Ok(())
}
fn validate_store_schema(conn: &Connection) -> rusqlite::Result<bool> {
let meta_columns = table_columns(conn, "stored_configs")?;
let field_columns = table_columns(conn, "stored_config_fields")?;
let required_meta = [
"config_id",
"display_name",
"created_at",
"updated_at",
"favorite",
"temporary",
];
let required_fields = ["config_id", "field_name", "field_json", "updated_at"];
Ok(required_meta
.iter()
.all(|column| meta_columns.contains(*column))
&& required_fields
.iter()
.all(|column| field_columns.contains(*column)))
}
fn move_db_file_if_exists(path: &Path) -> bool {
if !path.exists() {
return true;
}
let target = PathBuf::from(format!(
"{}.corrupt.{}",
path.to_string_lossy(),
now_ts_string()
));
match std::fs::rename(path, &target) {
Ok(_) => true,
Err(e) => {
ohrs_log_error!(
"[Rust] failed to move corrupt config db {} to {}: {}",
path.display(),
target.display(),
e
);
false
}
}
}
fn recover_config_db_files(path: &Path) -> bool {
let main_ok = move_db_file_if_exists(path);
let wal_ok = move_db_file_if_exists(Path::new(&format!("{}-wal", path.to_string_lossy())));
let shm_ok = move_db_file_if_exists(Path::new(&format!("{}-shm", path.to_string_lossy())));
main_ok && wal_ok && shm_ok
}
fn open_connection(path: &Path) -> Option<Connection> {
let conn = match Connection::open(path) {
Ok(conn) => conn,
Err(e) => {
hilog_error!("[Rust] failed to open config db {}: {}", path.display(), e);
ohrs_log_error!("[Rust] failed to open config db {}: {}", path.display(), e);
return None;
}
};
if let Err(e) = init_schema(&conn) {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to initialize config db {}: {}",
path.display(),
e
);
return None;
drop(conn);
if !recover_config_db_files(path) {
return None;
}
let recovered = match Connection::open(path) {
Ok(conn) => conn,
Err(e) => {
ohrs_log_error!(
"[Rust] failed to open recovered config db {}: {}",
path.display(),
e
);
return None;
}
};
if let Err(e) = init_schema(&recovered) {
ohrs_log_error!(
"[Rust] failed to initialize recovered config db {}: {}",
path.display(),
e
);
return None;
}
return Some(recovered);
}
Some(conn)
}
pub(crate) fn open_db() -> Option<ConfigDbGuard<'static>> {
let path = db_file_path()?;
let mut guard = match CONFIG_DB_CONNECTION.lock() {
Ok(guard) => guard,
Err(e) => {
ohrs_log_error!("[Rust] failed to lock config db connection: {}", e);
return None;
}
};
let should_open = guard
.as_ref()
.map(|cached| cached.path != path || !cached.path.exists())
.unwrap_or(true);
if should_open {
let conn = open_connection(&path)?;
*guard = Some(CachedConfigDb { path, conn });
}
Some(ConfigDbGuard { guard })
}
fn row_to_meta(row: &rusqlite::Row<'_>) -> rusqlite::Result<StoredConfigMetaRecord> {
Ok(StoredConfigMetaRecord {
config_id: row.get(0)?,
@@ -125,38 +330,64 @@ fn validate_snapshot_schema(conn: &Connection) -> bool {
has_stored_configs && has_stored_fields
}
fn copy_snapshot_tables(src: &Connection, dst: &mut Connection) -> rusqlite::Result<()> {
fn read_snapshot_tables(
src: &Connection,
) -> rusqlite::Result<(Vec<StoredConfigMetaRecord>, Vec<SnapshotFieldRow>)> {
src.execute_batch("BEGIN DEFERRED TRANSACTION")?;
let mut meta_rows = Vec::<StoredConfigMetaRecord>::new();
{
let mut stmt = src.prepare(
"SELECT config_id, display_name, created_at, updated_at, favorite, temporary
FROM stored_configs",
)?;
let rows = stmt.query_map([], row_to_meta)?;
for row in rows {
meta_rows.push(row?);
let mut field_rows = Vec::<SnapshotFieldRow>::new();
let read_result = (|| -> rusqlite::Result<()> {
{
let mut stmt = src.prepare(
"SELECT config_id, display_name, created_at, updated_at, favorite, temporary
FROM stored_configs",
)?;
let rows = stmt.query_map([], row_to_meta)?;
for row in rows {
meta_rows.push(row?);
}
}
{
let mut stmt = src.prepare(
"SELECT config_id, field_name, field_json, updated_at
FROM stored_config_fields",
)?;
let rows = stmt.query_map([], |row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, String>(3)?,
))
})?;
for row in rows {
field_rows.push(row?);
}
}
}
let mut field_rows = Vec::<(String, String, String, String)>::new();
{
let mut stmt = src.prepare(
"SELECT config_id, field_name, field_json, updated_at
FROM stored_config_fields",
)?;
let rows = stmt.query_map([], |row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, String>(3)?,
))
})?;
for row in rows {
field_rows.push(row?);
Ok(())
})();
match read_result {
Ok(()) => {
src.execute_batch("COMMIT")?;
Ok((meta_rows, field_rows))
}
Err(err) => {
let _ = src.execute_batch("ROLLBACK");
Err(err)
}
}
}
fn write_snapshot_tables(
dst: &mut Connection,
meta_rows: Vec<StoredConfigMetaRecord>,
field_rows: Vec<SnapshotFieldRow>,
) -> rusqlite::Result<()> {
let tx = dst.unchecked_transaction()?;
tx.execute("DELETE FROM stored_config_fields", [])?;
tx.execute("DELETE FROM stored_configs", [])?;
@@ -188,12 +419,17 @@ fn copy_snapshot_tables(src: &Connection, dst: &mut Connection) -> rusqlite::Res
tx.commit()
}
fn copy_snapshot_tables(src: &Connection, dst: &mut Connection) -> rusqlite::Result<()> {
let (meta_rows, field_rows) = read_snapshot_tables(src)?;
write_snapshot_tables(dst, meta_rows, field_rows)
}
fn ensure_parent_dir(path: &Path) -> bool {
match path.parent() {
Some(parent) => match std::fs::create_dir_all(parent) {
Ok(_) => true,
Err(e) => {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to create snapshot parent {}: {}",
parent.display(),
e
@@ -219,7 +455,7 @@ fn to_meta(record: StoredConfigMetaRecord) -> StoredConfigMeta {
pub fn init_config_meta_store(root_dir: String) -> bool {
let root = PathBuf::from(root_dir);
if let Err(e) = std::fs::create_dir_all(&root) {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to create config db dir {}: {}",
root.display(),
e
@@ -233,7 +469,7 @@ pub fn init_config_meta_store(root_dir: String) -> bool {
*guard = Some(db_path.clone());
}
Err(e) => {
hilog_error!("[Rust] failed to lock config db path: {}", e);
ohrs_log_error!("[Rust] failed to lock config db path: {}", e);
return false;
}
}
@@ -242,7 +478,7 @@ pub fn init_config_meta_store(root_dir: String) -> bool {
return false;
}
hilog_debug!("[Rust] initialized config db at {}", db_path.display());
ohrs_log_debug!("[Rust] initialized config db at {}", db_path.display());
true
}
@@ -257,7 +493,7 @@ pub fn export_config_store_snapshot(target_path: String) -> bool {
let mut dst = match Connection::open(&target) {
Ok(conn) => conn,
Err(e) => {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to open snapshot target {}: {}",
target.display(),
e
@@ -266,7 +502,7 @@ pub fn export_config_store_snapshot(target_path: String) -> bool {
}
};
if let Err(e) = init_schema(&dst) {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to init snapshot schema {}: {}",
target.display(),
e
@@ -276,7 +512,7 @@ pub fn export_config_store_snapshot(target_path: String) -> bool {
match copy_snapshot_tables(&src, &mut dst) {
Ok(_) => true,
Err(e) => {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to export snapshot {}: {}",
target.display(),
e
@@ -286,34 +522,92 @@ pub fn export_config_store_snapshot(target_path: String) -> bool {
}
}
pub fn import_config_store_snapshot(source_path: String) -> bool {
pub fn import_config_store_snapshot_with_result(source_path: String) -> SnapshotImportResult {
let source = PathBuf::from(source_path);
let src = match Connection::open(&source) {
Ok(conn) => conn,
Err(e) => {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to open snapshot source {}: {}",
source.display(),
e
);
return snapshot_import_err("source_open_failed", e.to_string(), false);
}
};
if !validate_snapshot_schema(&src) {
ohrs_log_error!("[Rust] invalid snapshot schema {}", source.display());
return snapshot_import_err(
"invalid_snapshot_schema",
format!("invalid snapshot schema: {}", source.display()),
true,
);
}
let (meta_rows, field_rows) = match read_snapshot_tables(&src) {
Ok(rows) => rows,
Err(e) => {
ohrs_log_error!(
"[Rust] failed to read snapshot source {}: {}",
source.display(),
e
);
return snapshot_import_err("invalid_snapshot_data", e.to_string(), true);
}
};
let Some(mut dst) = open_db() else {
return snapshot_import_err(
"destination_open_failed",
"failed to open local config store",
false,
);
};
match write_snapshot_tables(&mut dst, meta_rows, field_rows) {
Ok(_) => snapshot_import_ok(),
Err(e) => {
ohrs_log_error!(
"[Rust] failed to import snapshot {}: {}",
source.display(),
e
);
snapshot_import_err("destination_write_failed", e.to_string(), false)
}
}
}
pub fn import_config_store_snapshot(source_path: String) -> bool {
import_config_store_snapshot_with_result(source_path).ok
}
pub fn reset_config_meta_store() -> bool {
let Some(conn) = open_db() else {
return false;
};
let tx = match conn.unchecked_transaction() {
Ok(tx) => tx,
Err(e) => {
ohrs_log_error!(
"[Rust] failed to start config store reset transaction: {}",
e
);
return false;
}
};
if !validate_snapshot_schema(&src) {
hilog_error!("[Rust] invalid snapshot schema {}", source.display());
if let Err(e) = tx.execute("DELETE FROM stored_config_fields", []) {
ohrs_log_error!("[Rust] failed to reset config fields: {}", e);
let _ = tx.rollback();
return false;
}
let Some(mut dst) = open_db() else {
if let Err(e) = tx.execute("DELETE FROM stored_configs", []) {
ohrs_log_error!("[Rust] failed to reset config meta: {}", e);
let _ = tx.rollback();
return false;
};
match copy_snapshot_tables(&src, &mut dst) {
}
match tx.commit() {
Ok(_) => true,
Err(e) => {
hilog_error!(
"[Rust] failed to import snapshot {}: {}",
source.display(),
e
);
ohrs_log_error!("[Rust] failed to commit config store reset: {}", e);
false
}
}
@@ -331,7 +625,7 @@ pub fn list_config_meta_entries() -> StoredConfigList {
) {
Ok(stmt) => stmt,
Err(e) => {
hilog_error!("[Rust] failed to prepare list meta query: {}", e);
ohrs_log_error!("[Rust] failed to prepare list meta query: {}", e);
return StoredConfigList { configs: vec![] };
}
};
@@ -339,7 +633,7 @@ pub fn list_config_meta_entries() -> StoredConfigList {
let rows = match stmt.query_map([], row_to_meta) {
Ok(rows) => rows,
Err(e) => {
hilog_error!("[Rust] failed to list config meta rows: {}", e);
ohrs_log_error!("[Rust] failed to list config meta rows: {}", e);
return StoredConfigList { configs: vec![] };
}
};
@@ -358,59 +652,6 @@ pub fn get_config_meta(config_id: &str) -> Option<StoredConfigMeta> {
load_meta_record(&conn, config_id).map(to_meta)
}
pub fn upsert_config_meta(
config_id: String,
display_name: String,
favorite: bool,
temporary: bool,
) -> StoredConfigMeta {
let now = now_ts_string();
let Some(conn) = open_db() else {
return StoredConfigMeta {
config_id,
display_name,
created_at: now.clone(),
updated_at: now,
favorite,
temporary,
};
};
let created_at = load_meta_record(&conn, &config_id)
.map(|record| record.created_at)
.unwrap_or_else(|| now.clone());
if let Err(e) = conn.execute(
"INSERT INTO stored_configs (
config_id, display_name, created_at, updated_at, favorite, temporary
) VALUES (?1, ?2, ?3, ?4, ?5, ?6)
ON CONFLICT(config_id) DO UPDATE SET
display_name = excluded.display_name,
updated_at = excluded.updated_at,
favorite = excluded.favorite,
temporary = excluded.temporary",
params![
config_id,
display_name,
created_at,
now,
if favorite { 1 } else { 0 },
if temporary { 1 } else { 0 }
],
) {
hilog_error!("[Rust] failed to upsert config meta: {}", e);
}
get_config_meta(&config_id).unwrap_or(StoredConfigMeta {
config_id,
display_name,
created_at,
updated_at: now,
favorite,
temporary,
})
}
pub(crate) fn upsert_config_meta_in_tx(
tx: &rusqlite::Transaction<'_>,
config_id: String,
@@ -492,19 +733,45 @@ pub fn set_config_display_name(
Some(to_meta(record))
}
pub fn delete_config_meta(config_id: &str) -> bool {
let Some(conn) = open_db() else {
return false;
};
pub fn set_config_favorite(config_id: String, favorite: bool) -> Option<StoredConfigMeta> {
let conn = open_db()?;
let now = now_ts_string();
let tx = conn.unchecked_transaction().ok()?;
match conn.execute(
"DELETE FROM stored_configs WHERE config_id = ?1",
params![config_id],
) {
Ok(rows) => rows > 0,
Err(e) => {
hilog_error!("[Rust] failed to delete config meta {}: {}", config_id, e);
false
}
if favorite {
tx.execute(
"UPDATE stored_configs
SET favorite = 0,
updated_at = CASE WHEN favorite != 0 THEN ?1 ELSE updated_at END
WHERE favorite != 0 AND config_id <> ?2",
params![now, config_id.clone()],
)
.ok()?;
}
let rows = tx
.execute(
"UPDATE stored_configs
SET favorite = ?2, updated_at = ?3
WHERE config_id = ?1",
params![config_id.clone(), if favorite { 1 } else { 0 }, now],
)
.ok()?;
if rows == 0 {
return None;
}
let meta = tx
.query_row(
"SELECT config_id, display_name, created_at, updated_at, favorite, temporary
FROM stored_configs WHERE config_id = ?1",
params![config_id],
row_to_meta,
)
.optional()
.ok()
.flatten()
.map(to_meta)?;
tx.commit().ok()?;
Some(meta)
}
@@ -35,14 +35,6 @@ pub struct ExportTomlResult {
pub toml_text: String,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
#[napi(object)]
pub struct StoredConfigSummary {
pub config_id: String,
pub display_name: String,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
#[napi(object)]
@@ -66,3 +58,13 @@ pub struct KeyValuePair {
pub key: String,
pub value: String,
}
#[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,
}
+139 -31
View File
@@ -1,21 +1,74 @@
use super::{field_store, import_export, legacy_migration, validation};
use crate::config::storage::config_meta::{
delete_config_meta, get_config_meta, init_config_meta_store, list_config_meta_entries, open_db,
upsert_config_meta_in_tx,
get_config_meta, init_config_meta_store, list_config_meta_entries, open_db,
reset_config_meta_store, upsert_config_meta_in_tx,
};
use crate::config::types::stored_config::{ExportTomlResult, StoredConfigRecord};
use easytier::common::config::ConfigLoader;
use easytier::proto::api::manage::NetworkConfig;
use ohos_hilog_binding::{hilog_debug, hilog_error};
use once_cell::sync::Lazy;
use rusqlite::params;
use serde_json::Value;
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Mutex;
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";
#[derive(Clone)]
pub(crate) struct RuntimeConfigSnapshot {
pub display_name: String,
pub config: NetworkConfig,
}
pub(crate) fn cache_runtime_config_snapshot(
config_id: String,
display_name: String,
config: NetworkConfig,
) {
if let Ok(mut guard) = RUNTIME_CONFIG_SNAPSHOTS.lock() {
guard.insert(
config_id,
RuntimeConfigSnapshot {
display_name,
config,
},
);
}
}
pub(crate) 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> {
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> {
CONFIG_ROOT_DIR
.lock()
@@ -35,7 +88,7 @@ pub fn init_config_store(root_dir: String) -> bool {
let root = PathBuf::from(root_dir);
let configs_dir = root.join(CONFIG_DIR_NAME);
if let Err(e) = std::fs::create_dir_all(&configs_dir) {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to create config dir {}: {}",
configs_dir.display(),
e
@@ -48,7 +101,7 @@ pub fn init_config_store(root_dir: String) -> bool {
*guard = Some(root.clone());
}
Err(e) => {
hilog_error!("[Rust] failed to lock config root dir: {}", e);
ohrs_log_error!("[Rust] failed to lock config root dir: {}", e);
return false;
}
}
@@ -57,14 +110,27 @@ pub fn init_config_store(root_dir: String) -> bool {
return false;
}
hilog_debug!(
ohrs_log_debug!(
"[Rust] initialized config repo at {}",
configs_dir.display()
);
true
}
pub fn reset_config_store() -> bool {
if !reset_config_meta_store() {
return false;
}
if let Ok(mut guard) = RUNTIME_CONFIG_SNAPSHOTS.lock() {
guard.clear();
}
true
}
fn migrate_legacy_file_if_needed(config_id: &str) -> Option<()> {
if validation::validate_config_id(config_id).is_err() {
return None;
}
legacy_migration::migrate_legacy_file_if_needed(
&config_root_dir(),
CONFIG_DIR_NAME,
@@ -81,7 +147,7 @@ pub fn save_config_record(
let config = match validation::validate_config_json(&config_json, config_id.clone()) {
Ok(config) => config,
Err(e) => {
hilog_error!("[Rust] save_config_record failed {}", e);
ohrs_log_error!("[Rust] save_config_record failed {}", e);
return None;
}
};
@@ -89,7 +155,7 @@ pub fn save_config_record(
let normalized_json = match serde_json::to_string(&config) {
Ok(raw) => raw,
Err(e) => {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to serialize normalized config {}: {}",
config_id,
e
@@ -105,15 +171,15 @@ pub fn save_config_record(
let conn = open_db()?;
let tx = conn.unchecked_transaction().ok()?;
let existing_meta = get_config_meta(&config_id);
let favorite = existing_meta
.as_ref()
.map(|meta| meta.favorite)
.unwrap_or(false);
let temporary = existing_meta
.as_ref()
.map(|meta| meta.temporary)
.unwrap_or(false);
let existing_meta = tx
.query_row(
"SELECT favorite, temporary FROM stored_configs WHERE config_id = ?1",
params![config_id.clone()],
|row| Ok((row.get::<_, i64>(0)? != 0, row.get::<_, i64>(1)? != 0)),
)
.ok();
let favorite = existing_meta.map(|meta| meta.0).unwrap_or(false);
let temporary = existing_meta.map(|meta| meta.1).unwrap_or(false);
let meta = upsert_config_meta_in_tx(&tx, config_id.clone(), display_name, favorite, temporary)?;
field_store::replace_config_fields(&tx, &config_id, fields)?;
@@ -133,30 +199,52 @@ pub fn save_config_record(
}
pub fn load_config_json(config_id: &str) -> Option<String> {
validation::validate_config_id(config_id).ok()?;
migrate_legacy_file_if_needed(config_id)?;
let object = field_store::load_config_map_from_db(config_id)?;
serde_json::to_string(&Value::Object(object)).ok()
}
pub fn get_config_record(config_id: &str) -> Option<StoredConfigRecord> {
validation::validate_config_id(config_id).ok()?;
let config_json = load_config_json(config_id)?;
let meta = get_config_meta(config_id)?;
Some(StoredConfigRecord { meta, config_json })
}
pub fn get_config_field_value(config_id: &str, field: &str) -> Option<String> {
let total_start = Instant::now();
validation::validate_config_id(config_id).ok()?;
migrate_legacy_file_if_needed(config_id)?;
let open_start = Instant::now();
let conn = open_db()?;
conn.query_row(
"SELECT field_json FROM stored_config_fields
let open_elapsed = open_start.elapsed();
let query_start = Instant::now();
let result = conn
.query_row(
"SELECT field_json FROM stored_config_fields
WHERE config_id = ?1 AND field_name = ?2",
params![config_id, field],
|row| row.get::<_, String>(0),
)
.ok()
params![config_id, field],
|row| row.get::<_, String>(0),
)
.ok();
ohrs_log_debug!(
"[Rust] get_config_field_value config={} field={} found={} open_ms={} query_ms={} total_ms={} len={}",
config_id,
field,
result.is_some(),
open_elapsed.as_millis(),
query_start.elapsed().as_millis(),
total_start.elapsed().as_millis(),
result.as_ref().map(|value| value.len()).unwrap_or(0)
);
result
}
pub fn set_config_field_value(config_id: &str, field: &str, json_value: &str) -> bool {
if validation::validate_config_id(config_id).is_err() {
return false;
}
if field.contains('.') {
return false;
}
@@ -191,15 +279,12 @@ pub fn set_config_field_value(config_id: &str, field: &str, json_value: &str) ->
save_config_record(config_id.to_string(), display_name, normalized).is_some()
}
pub fn get_display_name(config_id: &str) -> Option<String> {
get_config_meta(config_id).map(|meta| meta.display_name)
}
pub fn get_default_config_json() -> Option<String> {
crate::build_default_network_config_json().ok()
}
pub fn create_config_record(config_id: String, display_name: String) -> Option<StoredConfigRecord> {
validation::validate_config_id(&config_id).ok()?;
let raw = get_default_config_json()?;
let mut config = serde_json::from_str::<NetworkConfig>(&raw).ok()?;
config.instance_id = Some(config_id.clone());
@@ -208,11 +293,21 @@ pub fn create_config_record(config_id: String, display_name: String) -> Option<S
}
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,
};
crate::run_network_instance_from_json(&raw)
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 {
@@ -220,6 +315,9 @@ pub fn list_config_meta_json() -> String {
}
pub fn delete_config_record(config_id: &str) -> bool {
if validation::validate_config_id(config_id).is_err() {
return false;
}
if let Some(path) = legacy_config_file_path(config_id) {
if path.exists() {
let _ = std::fs::remove_file(path);
@@ -234,14 +332,24 @@ pub fn delete_config_record(config_id: &str) -> bool {
"DELETE FROM stored_config_fields WHERE config_id = ?1",
params![config_id],
) {
hilog_error!("[Rust] failed to delete config fields {}: {}", config_id, e);
ohrs_log_error!("[Rust] failed to delete config fields {}: {}", config_id, e);
return false;
}
delete_config_meta(config_id)
match conn.execute(
"DELETE FROM stored_configs WHERE config_id = ?1",
params![config_id],
) {
Ok(rows) => rows > 0,
Err(e) => {
ohrs_log_error!("[Rust] failed to delete config meta {}: {}", config_id, e);
false
}
}
}
pub fn export_config_toml(config_id: &str) -> Option<ExportTomlResult> {
validation::validate_config_id(config_id).ok()?;
let record = get_config_record(config_id)?;
import_export::export_config_toml_from_record(&record)
}
@@ -1,5 +1,4 @@
use crate::config::storage::config_meta::{now_ts_string, open_db};
use ohos_hilog_binding::hilog_error;
use rusqlite::{Connection, params};
use serde_json::{Map, Value};
@@ -43,7 +42,7 @@ pub(super) fn replace_config_fields(
"DELETE FROM stored_config_fields WHERE config_id = ?1",
params![config_id],
) {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to clear existing config fields {}: {}",
config_id,
e
@@ -58,7 +57,7 @@ pub(super) fn replace_config_fields(
VALUES (?1, ?2, ?3, ?4)",
params![config_id, field_name, field_json, now_ts_string()],
) {
hilog_error!("[Rust] failed to persist config field {}: {}", config_id, e);
ohrs_log_error!("[Rust] failed to persist config field {}: {}", config_id, e);
return None;
}
}
@@ -1,12 +1,17 @@
use crate::config::storage::config_meta::get_config_meta;
use ohos_hilog_binding::hilog_error;
use std::path::PathBuf;
use super::validation;
pub(super) fn legacy_config_file_path(
root_dir: &Option<PathBuf>,
config_dir_name: &str,
config_id: &str,
) -> Option<PathBuf> {
if !validation::is_valid_config_id(config_id) {
ohrs_log_error!("[Rust] invalid legacy config_id {}", config_id);
return None;
}
root_dir.as_ref().map(|root| {
root.join(config_dir_name)
.join(format!("{}.json", config_id))
@@ -35,7 +40,7 @@ pub(super) fn migrate_legacy_file_if_needed(
save_config_record(config_id.to_string(), display_name, raw)?;
if let Err(e) = std::fs::remove_file(&legacy_path) {
hilog_error!(
ohrs_log_error!(
"[Rust] failed to remove legacy config file {}: {}",
legacy_path.display(),
e
@@ -1,13 +1,25 @@
use easytier::proto::api::manage::NetworkConfig;
use serde_json::{Map, Value};
use uuid::Uuid;
pub(super) fn validate_config_id(config_id: &str) -> Result<(), String> {
if config_id.is_empty() {
return Err("config_id is required".to_string());
}
Uuid::parse_str(config_id)
.map(|_| ())
.map_err(|e| format!("invalid config_id {}: {}", config_id, e))
}
pub(super) fn is_valid_config_id(config_id: &str) -> bool {
validate_config_id(config_id).is_ok()
}
pub(super) fn normalize_config_id(
mut config: NetworkConfig,
requested_id: String,
) -> Result<NetworkConfig, String> {
if requested_id.is_empty() {
return Err("config_id is required".to_string());
}
validate_config_id(&requested_id)?;
config.instance_id = Some(requested_id);
Ok(config)
}
@@ -1,9 +1,14 @@
use crate::config;
use crate::config::types::stored_config::SnapshotImportResult;
pub(crate) fn init_config_store(root_dir: String) -> bool {
config::repository::init_config_store(root_dir)
}
pub(crate) fn reset_config_store() -> bool {
config::repository::reset_config_store()
}
pub(crate) fn list_configs() -> String {
config::repository::list_config_meta_json()
}
@@ -36,6 +41,10 @@ pub(crate) fn set_config_field(config_id: String, field: String, json_value: Str
config::repository::set_config_field_value(&config_id, &field, &json_value)
}
pub(crate) fn set_config_favorite(config_id: String, favorite: bool) -> bool {
config::storage::config_meta::set_config_favorite(config_id, favorite).is_some()
}
pub(crate) fn import_toml(toml_text: String, display_name: Option<String>) -> Option<String> {
config::repository::import_toml_config(toml_text, display_name)
.map(|record| record.meta.config_id)
@@ -52,3 +61,9 @@ pub(crate) fn export_config_store_snapshot(target_path: String) -> bool {
pub(crate) fn import_config_store_snapshot(source_path: String) -> bool {
config::storage::config_meta::import_config_store_snapshot(source_path)
}
pub(crate) fn import_config_store_snapshot_with_result(
source_path: String,
) -> SnapshotImportResult {
config::storage::config_meta::import_config_store_snapshot_with_result(source_path)
}
@@ -1,18 +1,15 @@
use crate::config::repository::load_config_json;
use crate::config::storage::config_meta::get_config_display_name;
use crate::config::repository::{clear_runtime_config_snapshot, get_runtime_config_snapshot};
use crate::config::types::stored_config::KeyValuePair;
use crate::kernel_bridge::{
aggregate_requested_tun_routes, start_local_socket_server as start_local_socket_server_inner,
stop_local_socket_server as stop_local_socket_server_inner,
};
use crate::runtime::state::runtime_state::{
RuntimeAggregateState, TunAggregateState, clear_tun_attached, mark_tun_attached,
RuntimeAggregateState, RuntimeInstanceState, TunAggregateState, clear_tun_attached,
is_tun_attached, mark_tun_attached, runtime_instance_from_config_snapshot,
runtime_instance_from_running_info,
};
use crate::{ASYNC_RUNTIME, EASYTIER_VERSION, INSTANCE_MANAGER, WEB_CLIENTS};
use easytier::proto::api::manage::NetworkConfig;
use ohos_hilog_binding::{hilog_error, hilog_info};
use std::sync::Arc;
use crate::{ASYNC_RUNTIME, INSTANCE_MANAGER, WEB_CLIENTS};
pub(crate) fn start_kernel(
config_id: String,
@@ -29,9 +26,12 @@ pub(crate) fn stop_kernel(
) -> bool {
clear_tun_attached(&config_id);
if stop_web_client(&config_id) {
clear_runtime_config_snapshot(&config_id);
return true;
}
let _ = stop_local_socket_server_inner();
let Some(instance_id) = parse_instance_uuid(&config_id) else {
return false;
};
@@ -40,9 +40,20 @@ pub(crate) fn stop_kernel(
.delete_network_instance(vec![instance_id])
.map(|_| true)
.unwrap_or_else(|err| {
hilog_error!("[Rust] stop_kernel failed {}: {}", config_id, err);
ohrs_log_error!("[Rust] stop_kernel failed {}: {}", config_id, err);
false
});
if ret {
clear_runtime_config_snapshot(&config_id);
}
let has_active_instances = !INSTANCE_MANAGER.list_network_instance_ids().is_empty();
let has_web_clients = WEB_CLIENTS
.lock()
.map(|guard| !guard.is_empty())
.unwrap_or(false);
if has_active_instances || has_web_clients {
let _ = start_local_socket_server_inner();
}
maybe_stop_local_socket_server();
ret
}
@@ -59,10 +70,10 @@ pub(crate) fn stop_network_instance(
}
pub(crate) fn collect_network_infos() -> Vec<KeyValuePair> {
let infos = match INSTANCE_MANAGER.collect_network_infos_sync() {
let infos = match ASYNC_RUNTIME.block_on(INSTANCE_MANAGER.collect_network_infos()) {
Ok(infos) => infos,
Err(err) => {
hilog_error!("[Rust] collect network infos failed {}", err);
ohrs_log_error!("[Rust] collect network infos failed {}", err);
return vec![];
}
};
@@ -86,7 +97,7 @@ pub(crate) fn set_tun_fd(
parse_instance_uuid: impl Fn(&str) -> Option<uuid::Uuid>,
) -> bool {
let Some(instance_id) = parse_instance_uuid(&config_id) else {
hilog_error!("[Rust] set_tun_fd invalid instance id: {}", config_id);
ohrs_log_error!("[Rust] set_tun_fd invalid instance id: {}", config_id);
return false;
};
@@ -94,7 +105,7 @@ pub(crate) fn set_tun_fd(
.set_tun_fd(&instance_id, fd)
.map(|_| {
mark_tun_attached(&config_id);
hilog_info!(
ohrs_log_info!(
"[Rust] set_tun_fd success instance={} fd={} marked_attached=true",
config_id,
fd
@@ -102,20 +113,16 @@ pub(crate) fn set_tun_fd(
true
})
.unwrap_or_else(|err| {
hilog_error!("[Rust] set_tun_fd failed {}: {}", config_id, err);
ohrs_log_error!("[Rust] set_tun_fd failed {}: {}", config_id, err);
false
})
}
pub(crate) fn get_runtime_snapshot() -> RuntimeAggregateState {
get_runtime_snapshot_inner()
}
pub(crate) fn get_runtime_snapshot_inner() -> RuntimeAggregateState {
let infos = match INSTANCE_MANAGER.collect_network_infos_sync() {
pub(crate) fn collect_runtime_state() -> RuntimeAggregateState {
let infos = match ASYNC_RUNTIME.block_on(INSTANCE_MANAGER.collect_network_infos()) {
Ok(infos) => infos,
Err(err) => {
hilog_error!("[Rust] collect network infos failed {}", err);
ohrs_log_error!("[Rust] collect network infos failed {}", err);
return RuntimeAggregateState {
instances: vec![],
tun: TunAggregateState {
@@ -129,30 +136,67 @@ pub(crate) fn get_runtime_snapshot_inner() -> RuntimeAggregateState {
};
}
};
let mut live_infos = infos
.into_iter()
.map(|(instance_id, info)| (instance_id.to_string(), info))
.collect::<std::collections::HashMap<_, _>>();
let mut active_config_ids = live_infos.keys().cloned().collect::<Vec<_>>();
if let Ok(guard) = WEB_CLIENTS.lock() {
for config_id in guard.keys() {
if !active_config_ids.iter().any(|value| value == config_id) {
active_config_ids.push(config_id.clone());
}
}
}
let mut instances = Vec::with_capacity(infos.len());
for (instance_uuid, info) in infos {
let config_id = instance_uuid.to_string();
let display_name = get_config_display_name(&config_id).unwrap_or_else(|| config_id.clone());
let config_json = load_config_json(&config_id);
let stored_config = config_json
.as_deref()
.and_then(|raw| serde_json::from_str::<NetworkConfig>(raw).ok());
let magic_dns_enabled = stored_config
.as_ref()
.and_then(|cfg| cfg.enable_magic_dns)
.unwrap_or(false);
let need_exit_node = stored_config
.as_ref()
.map(|cfg| !cfg.exit_nodes.is_empty())
.unwrap_or(false);
instances.push(runtime_instance_from_running_info(
config_id,
display_name,
magic_dns_enabled,
need_exit_node,
info,
));
let mut instances = Vec::with_capacity(active_config_ids.len());
for config_id in active_config_ids {
if let Some(info) = live_infos.remove(&config_id) {
let snapshot = get_runtime_config_snapshot(&config_id);
let display_name = snapshot
.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,
info,
));
} else if let Some(snapshot) = get_runtime_config_snapshot(&config_id) {
instances.push(runtime_instance_from_config_snapshot(
config_id,
snapshot.display_name,
snapshot.config,
true,
));
} else {
let tun_attached = is_tun_attached(&config_id);
instances.push(RuntimeInstanceState {
config_id: config_id.clone(),
instance_id: config_id.clone(),
display_name: config_id.clone(),
running: true,
tun_required: tun_attached,
tun_attached,
magic_dns_enabled: false,
need_exit_node: false,
error_message: None,
my_node_info: None,
events: Vec::new(),
routes: Vec::new(),
peers: Vec::new(),
});
}
}
instances.sort_by(|a, b| {
@@ -32,6 +32,13 @@ pub(crate) fn send_local_socket_message(
Ok(())
}
fn shrink_clients_if_sparse(clients: &mut Vec<UnixStream>) {
let sparse_limit = clients.len().saturating_mul(2).max(4);
if clients.capacity() > sparse_limit {
clients.shrink_to_fit();
}
}
pub(crate) fn broadcast_local_socket_message(
clients: &mut Vec<UnixStream>,
message_type: &str,
@@ -45,6 +52,42 @@ pub(crate) fn broadcast_local_socket_message(
active_clients.push(client);
}
}
shrink_clients_if_sparse(&mut active_clients);
*clients = active_clients;
delivered
}
pub(crate) fn send_local_socket_json_payload_message(
stream: &mut UnixStream,
message_type: &str,
payload_json: &str,
) -> std::io::Result<()> {
let message_type_json = serde_json::to_string(message_type)
.map_err(|err| Error::new(ErrorKind::InvalidData, err.to_string()))?;
let mut raw = Vec::with_capacity(message_type_json.len() + payload_json.len() + 38);
raw.extend_from_slice(b"{\"messageType\":");
raw.extend_from_slice(message_type_json.as_bytes());
raw.extend_from_slice(b",\"payloadJson\":");
raw.extend_from_slice(payload_json.as_bytes());
raw.extend_from_slice(b"}\n");
stream.write_all(&raw)?;
Ok(())
}
pub(crate) fn broadcast_local_socket_json_payload_message(
clients: &mut Vec<UnixStream>,
message_type: &str,
payload_json: &str,
) -> bool {
let mut active_clients = Vec::with_capacity(clients.len());
let mut delivered = false;
for mut client in clients.drain(..) {
if send_local_socket_json_payload_message(&mut client, message_type, payload_json).is_ok() {
delivered = true;
active_clients.push(client);
}
}
shrink_clients_if_sparse(&mut active_clients);
*clients = active_clients;
delivered
}
@@ -1,20 +1,12 @@
use crate::config::repository::load_config_json;
use crate::config::repository::get_runtime_config_route_overrides;
use crate::runtime::state::runtime_state::RuntimeInstanceState;
use easytier::proto::api::manage::NetworkConfig;
use ipnet::IpNet;
use ohos_hilog_binding::hilog_debug;
use std::collections::HashSet;
use std::net::IpAddr;
pub(crate) fn load_manual_routes(config_id: &str) -> Vec<String> {
load_config_json(config_id)
.and_then(|raw| serde_json::from_str::<NetworkConfig>(&raw).ok())
.map(|config| config.routes)
.unwrap_or_default()
}
fn normalize_route_cidr(route: &str) -> Option<String> {
route
let normalized = route.split("->").next().unwrap_or(route).trim();
normalized
.parse::<IpNet>()
.ok()
.map(|network| match network {
@@ -22,7 +14,7 @@ fn normalize_route_cidr(route: &str) -> Option<String> {
IpNet::V6(net) => net.trunc().to_string(),
})
.or_else(|| {
route.parse::<IpAddr>().ok().map(|addr| match addr {
normalized.parse::<IpAddr>().ok().map(|addr| match addr {
IpAddr::V4(ip) => format!("{}/32", ip),
IpAddr::V6(ip) => format!("{}/128", ip),
})
@@ -67,8 +59,9 @@ pub(crate) fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec<Strin
.my_node_info
.as_ref()
.and_then(|info| info.virtual_ipv4_cidr.clone());
let manual_routes = load_manual_routes(&instance.config_id);
let proxy_cidrs = instance
let (manual_routes, config_proxy_cidrs) =
get_runtime_config_route_overrides(&instance.config_id);
let runtime_proxy_cidrs = instance
.routes
.iter()
.flat_map(|route| route.proxy_cidrs.iter().cloned())
@@ -80,15 +73,9 @@ pub(crate) fn aggregate_tun_routes(instance: &RuntimeInstanceState) -> Vec<Strin
}
raw_routes.extend(manual_routes.iter().cloned());
raw_routes.extend(proxy_cidrs.iter().cloned());
let aggregated_routes = simplify_routes(raw_routes);
hilog_debug!(
"[Rust] aggregate_tun_routes instance={} proxy_cidrs={:?} aggregated_routes={:?}",
instance.instance_id,
proxy_cidrs,
aggregated_routes
);
aggregated_routes
raw_routes.extend(config_proxy_cidrs.iter().cloned());
raw_routes.extend(runtime_proxy_cidrs.iter().cloned());
simplify_routes(raw_routes)
}
pub(crate) fn aggregate_requested_tun_routes(instances: &[RuntimeInstanceState]) -> Vec<String> {
@@ -1,17 +1,27 @@
use super::protocol::{TunRequestPayload, broadcast_local_socket_message};
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::get_runtime_snapshot_inner;
use crate::kernel_bridge::routing::aggregate_tun_routes;
use ohos_hilog_binding::{hilog_error, hilog_info};
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::proto::api::instance::ListPeerRequest;
use easytier::proto::rpc_types::controller::BaseController;
use once_cell::sync::Lazy;
use serde::Serialize;
use std::collections::{HashMap, HashSet};
use std::hash::Hash;
use std::io::ErrorKind;
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::PathBuf;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread::{self, JoinHandle};
use std::time::Duration;
use std::time::{Duration, Instant};
struct LocalSocketState {
stop_flag: std::sync::Arc<AtomicBool>,
@@ -20,12 +30,287 @@ struct LocalSocketState {
}
static LOCAL_SOCKET_STATE: Lazy<Mutex<Option<LocalSocketState>>> = Lazy::new(|| Mutex::new(None));
const SOCKET_TICK_INTERVAL: Duration = Duration::from_millis(250);
const TRAFFIC_STATS_INTERVAL: Duration = Duration::from_secs(1);
const INSTANCE_POLL_INTERVAL: Duration = Duration::from_secs(1);
const TUN_FAST_CHECK_WINDOW: Duration = Duration::from_secs(8);
const EVENT_RECEIVER_SYNC_INTERVAL: Duration = Duration::from_secs(1);
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
struct TrafficStatsPayload {
instances: Vec<InstanceTrafficStats>,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
struct InstanceTrafficStats {
config_id: String,
instance_id: String,
rx_bytes: i64,
tx_bytes: i64,
peers: Vec<PeerTrafficStats>,
}
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
struct PeerTrafficStats {
peer_id: i64,
rx_bytes: i64,
tx_bytes: i64,
total_bytes: i64,
latency_us: i64,
loss_rate: f64,
}
struct PendingPeerEvent {
event: &'static str,
instance_id: String,
peer_id: i64,
conn: Option<RuntimePeerConnInfo>,
}
#[derive(Default)]
struct DrainedKernelEvents {
tun_refresh: bool,
topology_lost: bool,
peer_events: Vec<PendingPeerEvent>,
}
#[derive(Serialize)]
#[serde(rename_all = "camelCase")]
struct RuntimePeerEventPayload {
event: &'static str,
config_id: String,
instance_id: String,
peer_id: i64,
conn: Option<RuntimePeerConnInfo>,
}
fn shrink_hash_map_if_sparse<K: Eq + Hash, V>(map: &mut HashMap<K, V>) {
let sparse_limit = map.len().saturating_mul(2).max(8);
if map.capacity() > sparse_limit {
map.shrink_to_fit();
}
}
fn shrink_hash_set_if_sparse<T: Eq + Hash>(set: &mut HashSet<T>) {
let sparse_limit = set.len().saturating_mul(2).max(8);
if set.capacity() > sparse_limit {
set.shrink_to_fit();
}
}
fn sync_tun_event_receivers(receivers: &mut HashMap<String, EventBusSubscriber>) {
let mut active_instance_ids = HashSet::new();
for instance in INSTANCE_MANAGER.iter() {
let instance_id = instance.key().to_string();
active_instance_ids.insert(instance_id.clone());
if !receivers.contains_key(&instance_id)
&& let Some(receiver) = instance.value().subscribe_event()
{
receivers.insert(instance_id, receiver);
}
}
receivers.retain(|instance_id, _| active_instance_ids.contains(instance_id));
shrink_hash_map_if_sparse(receivers);
}
fn event_needs_tun_refresh(event: &GlobalCtxEvent) -> bool {
matches!(
event,
GlobalCtxEvent::DhcpIpv4Changed(_, _)
| GlobalCtxEvent::ProxyCidrsUpdated(_, _)
| GlobalCtxEvent::PublicIpv6RoutesUpdated(_, _)
)
}
fn drain_kernel_events(receivers: &mut HashMap<String, EventBusSubscriber>) -> DrainedKernelEvents {
let mut drained = DrainedKernelEvents::default();
let mut closed_receivers = Vec::new();
for (instance_id, receiver) in receivers.iter_mut() {
loop {
match receiver.try_recv() {
Ok(event) => {
drained.tun_refresh = event_needs_tun_refresh(&event) || drained.tun_refresh;
match event {
GlobalCtxEvent::PeerAdded(peer_id) => {
drained.peer_events.push(PendingPeerEvent {
event: "peer_added",
instance_id: instance_id.clone(),
peer_id: peer_id as i64,
conn: None,
});
}
GlobalCtxEvent::PeerRemoved(peer_id) => {
drained.peer_events.push(PendingPeerEvent {
event: "peer_removed",
instance_id: instance_id.clone(),
peer_id: peer_id as i64,
conn: None,
});
}
GlobalCtxEvent::PeerConnAdded(conn_info) => {
let peer_id = conn_info.peer_id as i64;
drained.peer_events.push(PendingPeerEvent {
event: "peer_conn_added",
instance_id: instance_id.clone(),
peer_id,
conn: Some(peer_conn_to_view(conn_info)),
});
}
GlobalCtxEvent::PeerConnRemoved(conn_info) => {
let peer_id = conn_info.peer_id as i64;
drained.peer_events.push(PendingPeerEvent {
event: "peer_conn_removed",
instance_id: instance_id.clone(),
peer_id,
conn: Some(peer_conn_to_view(conn_info)),
});
}
_ => {}
}
}
Err(tokio::sync::broadcast::error::TryRecvError::Empty) => break,
Err(tokio::sync::broadcast::error::TryRecvError::Lagged(_)) => {
drained.topology_lost = true;
continue;
}
Err(tokio::sync::broadcast::error::TryRecvError::Closed) => {
closed_receivers.push(instance_id.clone());
break;
}
}
}
}
for instance_id in closed_receivers {
receivers.remove(&instance_id);
}
drained
}
fn broadcast_runtime_peer_events(
clients: &mut Vec<UnixStream>,
peer_events: Vec<PendingPeerEvent>,
) {
for event in peer_events {
let payload = RuntimePeerEventPayload {
event: event.event,
config_id: event.instance_id.clone(),
instance_id: event.instance_id,
peer_id: event.peer_id,
conn: event.conn,
};
match serde_json::to_string(&payload) {
Ok(json) => {
let _ = broadcast_local_socket_json_payload_message(
clients,
"runtime_peer_event",
&json,
);
}
Err(err) => {
ohrs_log_error!("[Rust] serialize runtime peer event failed: {}", err);
}
}
}
}
fn tun_candidate_ids(snapshot: &RuntimeAggregateState) -> HashSet<String> {
snapshot
.instances
.iter()
.filter(|instance| instance.running && instance.tun_required)
.map(|instance| instance.instance_id.clone())
.collect()
}
fn collect_traffic_stats() -> TrafficStatsPayload {
let services = INSTANCE_MANAGER
.iter()
.filter_map(|instance| {
instance
.value()
.get_api_service()
.map(|api_service| (instance.key().to_string(), api_service))
})
.collect::<Vec<_>>();
let instances = ASYNC_RUNTIME.block_on(async {
let mut instances = Vec::new();
for (instance_id, api_service) in services {
let peers = match api_service
.get_peer_manage_service()
.list_peer(BaseController::default(), ListPeerRequest::default())
.await
{
Ok(response) => response.peer_infos,
Err(err) => {
ohrs_log_debug!(
"[Rust] collect traffic stats list_peer failed instance={}: {}",
instance_id,
err
);
continue;
}
};
let mut instance_rx_bytes = 0i64;
let mut instance_tx_bytes = 0i64;
let mut peer_stats = Vec::with_capacity(peers.len());
for peer in peers {
let mut peer_rx_bytes = 0i64;
let mut peer_tx_bytes = 0i64;
let mut latency_us = i64::MAX;
let mut loss_rate = 0f64;
for conn in peer.conns {
if let Some(stats) = conn.stats {
let rx_bytes = stats.rx_bytes as i64;
let tx_bytes = stats.tx_bytes as i64;
peer_rx_bytes += rx_bytes;
peer_tx_bytes += tx_bytes;
latency_us = latency_us.min(stats.latency_us as i64);
}
loss_rate = loss_rate.max(conn.loss_rate as f64);
}
instance_rx_bytes += peer_rx_bytes;
instance_tx_bytes += peer_tx_bytes;
peer_stats.push(PeerTrafficStats {
peer_id: peer.peer_id as i64,
rx_bytes: peer_rx_bytes,
tx_bytes: peer_tx_bytes,
total_bytes: peer_rx_bytes + peer_tx_bytes,
latency_us: if latency_us == i64::MAX {
-1
} else {
latency_us
},
loss_rate,
});
}
instances.push(InstanceTrafficStats {
config_id: instance_id.clone(),
instance_id,
rx_bytes: instance_rx_bytes,
tx_bytes: instance_tx_bytes,
peers: peer_stats,
});
}
instances
});
TrafficStatsPayload { instances }
}
pub fn start_local_socket_server() -> bool {
let socket_path = match kernel_socket_path() {
Some(path) => path,
None => {
hilog_error!("[Rust] kernel socket path unavailable");
ohrs_log_error!("[Rust] kernel socket path unavailable");
return false;
}
};
@@ -34,7 +319,7 @@ pub fn start_local_socket_server() -> bool {
Ok(guard) if guard.is_some() => return true,
Ok(_) => {}
Err(err) => {
hilog_error!("[Rust] lock localsocket state failed: {}", err);
ohrs_log_error!("[Rust] lock localsocket state failed: {}", err);
return false;
}
}
@@ -46,7 +331,7 @@ pub fn start_local_socket_server() -> bool {
let listener = match UnixListener::bind(&socket_path) {
Ok(listener) => listener,
Err(err) => {
hilog_error!(
ohrs_log_error!(
"[Rust] bind localsocket failed {}: {}",
socket_path.display(),
err
@@ -55,7 +340,7 @@ pub fn start_local_socket_server() -> bool {
}
};
if let Err(err) = listener.set_nonblocking(true) {
hilog_error!("[Rust] set localsocket nonblocking failed: {}", err);
ohrs_log_error!("[Rust] set localsocket nonblocking failed: {}", err);
let _ = std::fs::remove_file(&socket_path);
return false;
}
@@ -63,102 +348,208 @@ pub fn start_local_socket_server() -> bool {
let stop_flag = std::sync::Arc::new(AtomicBool::new(false));
let worker_stop_flag = stop_flag.clone();
let worker = thread::spawn(move || {
let mut last_snapshot_json = String::new();
let mut last_topology_json = String::new();
let mut delivered_tun_requests = HashSet::new();
let mut last_tun_route_signatures = HashMap::<String, String>::new();
let mut tun_fast_until = Instant::now() + TUN_FAST_CHECK_WINDOW;
let mut tun_bootstrap_done = false;
let mut last_event_receiver_sync_at: Option<Instant> = None;
let mut last_traffic_stats_at: Option<Instant> = None;
let mut last_instance_poll_at: Option<Instant> = None;
let mut tun_event_receivers = HashMap::<String, EventBusSubscriber>::new();
let mut clients = Vec::<UnixStream>::new();
while !worker_stop_flag.load(Ordering::Relaxed) {
let mut full_topology_dirty = false;
let mut accepted_client = false;
loop {
match listener.accept() {
Ok((stream, _addr)) => {
accepted_client = true;
full_topology_dirty = true;
clients.push(stream);
tun_fast_until = Instant::now() + TUN_FAST_CHECK_WINDOW;
tun_bootstrap_done = false;
}
Err(err) if err.kind() == ErrorKind::WouldBlock => break,
Err(err) => {
hilog_error!("[Rust] accept localsocket failed: {}", err);
ohrs_log_error!("[Rust] accept localsocket failed: {}", err);
break;
}
}
}
let snapshot = get_runtime_snapshot_inner();
let snapshot_json = match serde_json::to_string(&snapshot) {
Ok(json) => json,
if clients.is_empty() {
if !last_topology_json.is_empty() {
last_topology_json.clear();
last_topology_json.shrink_to_fit();
}
delivered_tun_requests.clear();
shrink_hash_set_if_sparse(&mut delivered_tun_requests);
last_tun_route_signatures.clear();
shrink_hash_map_if_sparse(&mut last_tun_route_signatures);
tun_event_receivers.clear();
shrink_hash_map_if_sparse(&mut tun_event_receivers);
clients.shrink_to_fit();
last_event_receiver_sync_at = None;
last_traffic_stats_at = None;
last_instance_poll_at = None;
tun_bootstrap_done = false;
thread::sleep(SOCKET_TICK_INTERVAL);
continue;
}
let now = Instant::now();
let should_sync_event_receivers = accepted_client
|| last_event_receiver_sync_at
.map(|last| now.duration_since(last) >= EVENT_RECEIVER_SYNC_INTERVAL)
.unwrap_or(true);
if should_sync_event_receivers {
sync_tun_event_receivers(&mut tun_event_receivers);
last_event_receiver_sync_at = Some(now);
}
let drained_events = drain_kernel_events(&mut tun_event_receivers);
let tun_refresh = drained_events.tun_refresh;
let topology_lost = drained_events.topology_lost;
let peer_events = drained_events.peer_events;
if topology_lost {
full_topology_dirty = true;
}
if tun_refresh {
tun_bootstrap_done = false;
tun_fast_until = now + TUN_FAST_CHECK_WINDOW;
}
if !peer_events.is_empty() {
broadcast_runtime_peer_events(&mut clients, peer_events);
}
let should_collect_traffic_stats = last_traffic_stats_at
.map(|last| now.duration_since(last) >= TRAFFIC_STATS_INTERVAL)
.unwrap_or(true);
if should_collect_traffic_stats {
last_traffic_stats_at = Some(now);
match serde_json::to_string(&collect_traffic_stats()) {
Ok(json) => {
let _ = broadcast_local_socket_json_payload_message(
&mut clients,
"traffic_stats",
&json,
);
}
Err(err) => {
ohrs_log_error!("[Rust] serialize traffic stats failed: {}", err);
}
}
}
let should_poll_instance = last_instance_poll_at
.map(|last| now.duration_since(last) >= INSTANCE_POLL_INTERVAL)
.unwrap_or(true);
let should_collect_topology = accepted_client
|| full_topology_dirty
|| tun_refresh
|| should_poll_instance
|| (!tun_bootstrap_done && now < tun_fast_until);
if !should_collect_topology {
thread::sleep(SOCKET_TICK_INTERVAL);
continue;
}
let snapshot = collect_runtime_state_inner();
last_instance_poll_at = Some(now);
match serde_json::to_string(&snapshot) {
Ok(json) => {
if accepted_client || full_topology_dirty || json != last_topology_json {
let _ = broadcast_local_socket_json_payload_message(
&mut clients,
"runtime_topology",
&json,
);
last_topology_json = json;
}
}
Err(err) => {
hilog_error!("[Rust] serialize runtime snapshot failed: {}", err);
thread::sleep(Duration::from_millis(250));
ohrs_log_error!("[Rust] serialize runtime topology failed: {}", err);
}
}
let active_tun_candidate_ids = tun_candidate_ids(&snapshot);
delivered_tun_requests
.retain(|instance_id| active_tun_candidate_ids.contains(instance_id));
last_tun_route_signatures
.retain(|instance_id, _| active_tun_candidate_ids.contains(instance_id));
shrink_hash_set_if_sparse(&mut delivered_tun_requests);
shrink_hash_map_if_sparse(&mut last_tun_route_signatures);
let mut saw_running_instance = false;
let mut saw_tun_candidate = false;
for instance in snapshot.instances.iter() {
if instance.running {
saw_running_instance = true;
}
if !(instance.running && instance.tun_required) {
continue;
}
};
if accepted_client || snapshot_json != last_snapshot_json {
let _ = broadcast_local_socket_message(
&mut clients,
"runtime_snapshot",
&snapshot_json,
);
last_snapshot_json = snapshot_json;
}
for instance in snapshot.instances.iter() {
if instance.running && instance.tun_required {
let virtual_ipv4 = instance
.my_node_info
.as_ref()
.and_then(|info| info.virtual_ipv4.clone());
let virtual_ipv4_cidr = instance
.my_node_info
.as_ref()
.and_then(|info| info.virtual_ipv4_cidr.clone());
if clients.is_empty() {
saw_tun_candidate = true;
let virtual_ipv4 = instance
.my_node_info
.as_ref()
.and_then(|info| info.virtual_ipv4.clone());
let virtual_ipv4_cidr = instance
.my_node_info
.as_ref()
.and_then(|info| info.virtual_ipv4_cidr.clone());
if clients.is_empty() {
continue;
}
if virtual_ipv4.is_none() || virtual_ipv4_cidr.is_none() {
continue;
}
let aggregated_routes = aggregate_tun_routes(instance);
let route_signature = serde_json::to_string(&(
&virtual_ipv4,
&virtual_ipv4_cidr,
&aggregated_routes,
instance.magic_dns_enabled,
instance.need_exit_node,
))
.unwrap_or_else(|_| "[]".to_string());
let should_send = !delivered_tun_requests.contains(&instance.instance_id)
|| last_tun_route_signatures
.get(&instance.instance_id)
.map(|value| value != &route_signature)
.unwrap_or(true);
if !should_send {
continue;
}
let payload = TunRequestPayload {
config_id: instance.config_id.clone(),
instance_id: instance.instance_id.clone(),
display_name: instance.display_name.clone(),
virtual_ipv4,
virtual_ipv4_cidr,
aggregated_routes,
magic_dns_enabled: instance.magic_dns_enabled,
need_exit_node: instance.need_exit_node,
};
let payload_json = match serde_json::to_string(&payload) {
Ok(json) => json,
Err(err) => {
ohrs_log_error!("[Rust] serialize tun request failed: {}", err);
continue;
}
if virtual_ipv4.is_none() || virtual_ipv4_cidr.is_none() {
continue;
}
let aggregated_routes = aggregate_tun_routes(instance);
let route_signature = serde_json::to_string(&aggregated_routes)
.unwrap_or_else(|_| "[]".to_string());
let should_send = !delivered_tun_requests.contains(&instance.instance_id)
|| last_tun_route_signatures
.get(&instance.instance_id)
.map(|value| value != &route_signature)
.unwrap_or(true);
if !should_send {
continue;
}
let payload = TunRequestPayload {
config_id: instance.config_id.clone(),
instance_id: instance.instance_id.clone(),
display_name: instance.display_name.clone(),
virtual_ipv4,
virtual_ipv4_cidr,
aggregated_routes,
magic_dns_enabled: instance.magic_dns_enabled,
need_exit_node: instance.need_exit_node,
};
let payload_json = match serde_json::to_string(&payload) {
Ok(json) => json,
Err(err) => {
hilog_error!("[Rust] serialize tun request failed: {}", err);
continue;
}
};
if broadcast_local_socket_message(&mut clients, "tun_request", &payload_json) {
delivered_tun_requests.insert(instance.instance_id.clone());
last_tun_route_signatures
.insert(instance.instance_id.clone(), route_signature);
}
} else {
delivered_tun_requests.remove(&instance.instance_id);
last_tun_route_signatures.remove(&instance.instance_id);
};
if broadcast_local_socket_message(&mut clients, "tun_request", &payload_json) {
delivered_tun_requests.insert(instance.instance_id.clone());
last_tun_route_signatures.insert(instance.instance_id.clone(), route_signature);
}
}
if !delivered_tun_requests.is_empty()
|| (saw_running_instance && !saw_tun_candidate)
|| now >= tun_fast_until
{
tun_bootstrap_done = true;
}
thread::sleep(Duration::from_millis(250));
thread::sleep(SOCKET_TICK_INTERVAL);
}
});
@@ -172,7 +563,7 @@ pub fn start_local_socket_server() -> bool {
true
}
Err(err) => {
hilog_error!("[Rust] lock localsocket state failed: {}", err);
ohrs_log_error!("[Rust] lock localsocket state failed: {}", err);
false
}
}
@@ -182,7 +573,7 @@ pub fn stop_local_socket_server() -> bool {
let state = match LOCAL_SOCKET_STATE.lock() {
Ok(mut guard) => guard.take(),
Err(err) => {
hilog_error!("[Rust] lock localsocket state failed: {}", err);
ohrs_log_error!("[Rust] lock localsocket state failed: {}", err);
return false;
}
};
+79 -34
View File
@@ -1,14 +1,46 @@
macro_rules! ohrs_log_error {
($($arg:tt)*) => {{
if $crate::platform::logging::log_manager::app_log_enabled(5) {
$crate::platform::logging::log_manager::record_app_log(
5,
"RustOhrs",
&std::format!($($arg)*),
);
}
}};
}
macro_rules! ohrs_log_info {
($($arg:tt)*) => {{
if $crate::platform::logging::log_manager::app_log_enabled(4) {
$crate::platform::logging::log_manager::record_app_log(
4,
"RustOhrs",
&std::format!($($arg)*),
);
}
}};
}
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 platform;
mod runtime;
use config::repository::{
create_config_record, delete_config_record, export_config_toml, get_config_field_value,
get_default_config_json, import_toml_config, init_config_store as init_repo_store,
list_config_meta_json, save_config_record, set_config_field_value, start_kernel_with_config_id,
};
use config::repository::{cache_runtime_config_snapshot, start_kernel_with_config_id};
use config::services::schema_service::{
ConfigFieldMapping, NetworkConfigSchema,
get_network_config_field_mappings as build_network_config_field_mappings,
@@ -20,7 +52,7 @@ 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};
use config::types::stored_config::{KeyValuePair, SharedConfigLinkPayload, SnapshotImportResult};
use easytier::common::constants::EASYTIER_VERSION;
use easytier::common::{
MachineIdOptions,
@@ -31,15 +63,11 @@ use easytier::proto::api::manage::NetworkConfig;
use easytier::proto::api::manage::NetworkingMethod;
use easytier::web_client::{WebClient, WebClientHooks, run_web_client};
use kernel_bridge::{
aggregate_requested_tun_routes, start_local_socket_server as start_local_socket_server_inner,
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 ohos_hilog_binding::{hilog_error, hilog_info};
use runtime::state::runtime_state::{
RuntimeAggregateState, TunAggregateState, clear_tun_attached, mark_tun_attached,
runtime_instance_from_running_info,
};
use runtime::state::runtime_state::RuntimeAggregateState;
use std::collections::{HashMap, HashSet};
use std::format;
use std::sync::{Arc, Mutex};
@@ -101,7 +129,7 @@ fn stop_web_client(config_id: &str) -> bool {
let managed = match WEB_CLIENTS.lock() {
Ok(mut guard) => guard.remove(config_id),
Err(err) => {
hilog_error!("[Rust] stop_web_client lock failed {}", err);
ohrs_log_error!("[Rust] stop_web_client lock failed {}", err);
return false;
}
};
@@ -127,7 +155,7 @@ fn stop_web_client(config_id: &str) -> bool {
.delete_network_instance(tracked_ids)
.map(|_| true)
.unwrap_or_else(|err| {
hilog_error!(
ohrs_log_error!(
"[Rust] stop config server instances failed {}: {}",
config_id,
err
@@ -160,12 +188,12 @@ fn run_config_server_instance(config_id: &str, config: &NetworkConfig) -> bool {
.next()
.is_some()
{
hilog_error!("[Rust] there is a running instance!");
ohrs_log_error!("[Rust] there is a running instance!");
return false;
}
let Some(config_server_url) = config.public_server_url.clone() else {
hilog_error!("[Rust] public_server_url missing for config server mode");
ohrs_log_error!("[Rust] public_server_url missing for config server mode");
return false;
};
let hooks = Arc::new(TrackedWebClientHooks::default());
@@ -192,7 +220,7 @@ fn run_config_server_instance(config_id: &str, config: &NetworkConfig) -> bool {
let client = match client {
Ok(client) => client,
Err(err) => {
hilog_error!("[Rust] start config server failed {}", err);
ohrs_log_error!("[Rust] start config server failed {}", err);
return false;
}
};
@@ -209,7 +237,7 @@ fn run_config_server_instance(config_id: &str, config: &NetworkConfig) -> bool {
true
}
Err(err) => {
hilog_error!("[Rust] store config server client failed {}", err);
ohrs_log_error!("[Rust] store config server client failed {}", err);
false
}
}
@@ -240,29 +268,33 @@ pub(crate) fn run_network_instance_from_json(cfg_json: &str) -> bool {
let config = match serde_json::from_str::<NetworkConfig>(cfg_json) {
Ok(cfg) => cfg,
Err(e) => {
hilog_error!("[Rust] parse config failed {}", e);
ohrs_log_error!("[Rust] parse config failed {}", e);
return false;
}
};
if is_config_server_config(&config) {
let Some(config_id) = config.instance_id.as_deref() else {
hilog_error!("[Rust] config server config missing instance id");
ohrs_log_error!("[Rust] config server config missing instance id");
return false;
};
return run_config_server_instance(config_id, &config);
let started = run_config_server_instance(config_id, &config);
if started {
cache_runtime_config_snapshot(config_id.to_string(), config_id.to_string(), config);
}
return started;
}
let cfg = match config.gen_config() {
Ok(toml) => toml,
Err(e) => {
hilog_error!("[Rust] parse config failed {}", e);
ohrs_log_error!("[Rust] parse config failed {}", e);
return false;
}
};
if !INSTANCE_MANAGER.list_network_instance_ids().is_empty() {
hilog_error!("[Rust] there is a running instance!");
ohrs_log_error!("[Rust] there is a running instance!");
return false;
}
@@ -275,14 +307,17 @@ pub(crate) fn run_network_instance_from_json(cfg_json: &str) -> bool {
.list_network_instance_ids()
.contains(&inst_id)
{
hilog_error!("[Rust] instance {} already exists", inst_id);
ohrs_log_error!("[Rust] instance {} already exists", inst_id);
return false;
}
match INSTANCE_MANAGER.run_network_instance(cfg, false, ConfigFileControl::STATIC_CONFIG) {
Ok(_) => true,
Ok(_) => {
cache_runtime_config_snapshot(inst_id.to_string(), inst_id.to_string(), config);
true
}
Err(err) => {
hilog_error!("[Rust] start_kernel failed for {}: {}", inst_id, err);
ohrs_log_error!("[Rust] start_kernel failed for {}: {}", inst_id, err);
false
}
}
@@ -292,7 +327,7 @@ fn parse_instance_uuid(config_id: &str) -> Option<Uuid> {
match Uuid::parse_str(config_id) {
Ok(uuid) => Some(uuid),
Err(err) => {
hilog_error!("[Rust] invalid config_id {}: {}", config_id, err);
ohrs_log_error!("[Rust] invalid config_id {}: {}", config_id, err);
None
}
}
@@ -303,6 +338,11 @@ pub fn init_config_store(root_dir: String) -> bool {
exports::config_api::init_config_store(root_dir)
}
#[napi]
pub fn reset_config_store() -> bool {
exports::config_api::reset_config_store()
}
#[napi]
pub fn list_configs() -> String {
exports::config_api::list_configs()
@@ -353,6 +393,11 @@ pub fn set_config_field(config_id: String, field: String, json_value: String) ->
exports::config_api::set_config_field(config_id, field, json_value)
}
#[napi]
pub fn set_config_favorite(config_id: String, favorite: bool) -> bool {
exports::config_api::set_config_favorite(config_id, favorite)
}
#[napi]
pub fn import_toml(toml_text: String, display_name: Option<String>) -> Option<String> {
exports::config_api::import_toml(toml_text, display_name)
@@ -373,6 +418,11 @@ pub fn import_config_store_snapshot(source_path: String) -> bool {
exports::config_api::import_config_store_snapshot(source_path)
}
#[napi]
pub fn import_config_store_snapshot_with_result(source_path: String) -> SnapshotImportResult {
exports::config_api::import_config_store_snapshot_with_result(source_path)
}
#[napi]
pub fn start_kernel(config_id: String) -> bool {
exports::runtime_api::start_kernel(config_id, start_kernel_with_config_id)
@@ -467,13 +517,8 @@ mod tests {
}
}
#[napi]
pub fn get_runtime_snapshot() -> RuntimeAggregateState {
exports::runtime_api::get_runtime_snapshot()
}
pub(crate) fn get_runtime_snapshot_inner() -> RuntimeAggregateState {
exports::runtime_api::get_runtime_snapshot_inner()
pub(crate) fn collect_runtime_state_inner() -> RuntimeAggregateState {
exports::runtime_api::collect_runtime_state()
}
#[napi]
@@ -0,0 +1,393 @@
use napi_derive_ohos::napi;
use once_cell::sync::Lazy;
use std::collections::VecDeque;
use std::fs::{self, Metadata, OpenOptions};
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
const LOG_DIR_NAME: &str = "easytier-logs";
const LOG_FILE_PREFIX: &str = "easytier-";
const LOG_FILE_SUFFIX: &str = ".log";
const MAX_LOG_FILES: usize = 10;
const MAX_MEMORY_LINES: usize = 500;
#[derive(Debug, Clone)]
#[napi(object)]
pub struct LogFileInfo {
pub file_name: String,
pub display_name: String,
pub size_bytes: i64,
pub modified_ms: i64,
pub active: bool,
}
#[derive(Clone)]
struct LogOptions {
core_log: bool,
debug_log: bool,
}
impl Default for LogOptions {
fn default() -> Self {
Self {
core_log: false,
debug_log: false,
}
}
}
#[derive(Default)]
struct LogManagerState {
log_dir: Option<PathBuf>,
active_file: Option<PathBuf>,
lines: VecDeque<String>,
options: LogOptions,
}
static LOG_MANAGER: Lazy<Mutex<LogManagerState>> =
Lazy::new(|| Mutex::new(LogManagerState::default()));
static CORE_LOG_ENABLED: AtomicBool = AtomicBool::new(false);
static DEBUG_LOG_ENABLED: AtomicBool = AtomicBool::new(false);
fn now_millis() -> u128 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_millis())
.unwrap_or(0)
}
fn sanitize_name(raw: &str) -> String {
let value = raw
.chars()
.map(|ch| {
if ch.is_ascii_alphanumeric() || ch == '-' || ch == '_' {
ch
} else {
'-'
}
})
.collect::<String>();
if value.is_empty() {
"process".to_string()
} else {
value
}
}
fn log_dir(root_dir: &str) -> PathBuf {
Path::new(root_dir).join(LOG_DIR_NAME)
}
fn is_log_file(path: &Path) -> bool {
path.file_name()
.and_then(|name| name.to_str())
.map(|name| name.starts_with(LOG_FILE_PREFIX) && name.ends_with(LOG_FILE_SUFFIX))
.unwrap_or(false)
}
fn sorted_log_files(dir: &Path) -> Vec<PathBuf> {
let mut files = fs::read_dir(dir)
.ok()
.into_iter()
.flat_map(|entries| entries.filter_map(|entry| entry.ok()))
.map(|entry| entry.path())
.filter(|path| is_log_file(path))
.collect::<Vec<_>>();
files.sort_by(|left, right| {
left.file_name()
.and_then(|name| name.to_str())
.unwrap_or_default()
.cmp(
right
.file_name()
.and_then(|name| name.to_str())
.unwrap_or_default(),
)
});
files
}
fn current_log_state() -> Option<(PathBuf, Option<PathBuf>)> {
LOG_MANAGER.lock().ok().and_then(|guard| {
guard
.log_dir
.clone()
.map(|dir| (dir, guard.active_file.clone()))
})
}
fn file_name(path: &Path) -> Option<String> {
path.file_name()
.and_then(|value| value.to_str())
.map(|value| value.to_string())
}
fn latest_process_log_file(dir: &Path, process_name: &str) -> Option<PathBuf> {
let suffix = format!("-{}{}", sanitize_name(process_name), LOG_FILE_SUFFIX);
sorted_log_files(dir).into_iter().rev().find(|path| {
path.file_name()
.and_then(|value| value.to_str())
.map(|value| value.ends_with(&suffix))
.unwrap_or(false)
})
}
fn modified_millis(metadata: &Metadata) -> i64 {
metadata
.modified()
.ok()
.and_then(|time| time.duration_since(UNIX_EPOCH).ok())
.map(|duration| duration.as_millis().min(i64::MAX as u128) as i64)
.unwrap_or(0)
}
fn resolve_log_file(dir: &Path, requested_name: &str) -> Option<PathBuf> {
if requested_name.contains('/')
|| requested_name.contains('\\')
|| requested_name.contains("..")
{
return None;
}
sorted_log_files(dir).into_iter().find(|path| {
path.file_name()
.and_then(|value| value.to_str())
.map(|value| value == requested_name)
.unwrap_or(false)
})
}
fn cleanup_old_logs(dir: &Path) {
let files = sorted_log_files(dir);
let overflow = files.len().saturating_sub(MAX_LOG_FILES);
for path in files.into_iter().take(overflow) {
let _ = fs::remove_file(path);
}
}
fn push_memory_line(state: &mut LogManagerState, line: String) {
state.lines.push_back(line);
while state.lines.len() > MAX_MEMORY_LINES {
state.lines.pop_front();
}
}
fn append_log_file(path: &Path, line: &str) {
if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) {
let _ = writeln!(file, "{}", line);
}
}
fn should_record_debug(level: i32) -> bool {
level <= 3
}
fn format_line(level: i32, target: &str, message: &str) -> String {
format!("{}[{}] {}", level, target, message.replace('\n', "\\n"))
}
pub(crate) fn configure(core_log: bool, debug_log: bool) {
CORE_LOG_ENABLED.store(core_log, Ordering::Relaxed);
DEBUG_LOG_ENABLED.store(debug_log, Ordering::Relaxed);
if let Ok(mut guard) = LOG_MANAGER.lock() {
guard.options.core_log = core_log;
guard.options.debug_log = debug_log;
}
}
pub(crate) fn app_log_enabled(level: i32) -> bool {
!should_record_debug(level) || DEBUG_LOG_ENABLED.load(Ordering::Relaxed)
}
pub(crate) fn core_log_enabled(level: i32) -> bool {
CORE_LOG_ENABLED.load(Ordering::Relaxed) && app_log_enabled(level)
}
pub(crate) fn record_app_log(level: i32, target: &str, message: &str) {
if !app_log_enabled(level) {
return;
}
if let Ok(mut guard) = LOG_MANAGER.lock() {
let line = format_line(level, target, message);
if let Some(path) = guard.active_file.as_ref() {
append_log_file(path, &line);
}
push_memory_line(&mut guard, line);
}
}
pub(crate) fn record_core_log(level: i32, target: &str, message: &str) {
if !core_log_enabled(level) {
return;
}
if let Ok(mut guard) = LOG_MANAGER.lock() {
let line = format_line(level, target, message);
if let Some(path) = guard.active_file.as_ref() {
append_log_file(path, &line);
}
push_memory_line(&mut guard, line);
}
}
#[napi]
pub fn init_log_manager(root_dir: String, process_name: String) -> bool {
let dir = log_dir(&root_dir);
if fs::create_dir_all(&dir).is_err() {
return false;
}
if LOG_MANAGER
.lock()
.map(|guard| guard.active_file.is_some())
.unwrap_or(false)
{
cleanup_old_logs(&dir);
return true;
}
let sanitized_process_name = sanitize_name(&process_name);
let active_file = if sanitized_process_name == "ui" {
dir.join(format!(
"{}{}-{}-{}{}",
LOG_FILE_PREFIX,
now_millis(),
std::process::id(),
sanitized_process_name,
LOG_FILE_SUFFIX
))
} else if let Some(path) = latest_process_log_file(&dir, "ui") {
path
} else {
dir.join(format!(
"{}{}-{}-{}{}",
LOG_FILE_PREFIX,
now_millis(),
std::process::id(),
sanitized_process_name,
LOG_FILE_SUFFIX
))
};
if OpenOptions::new()
.create(true)
.append(true)
.open(&active_file)
.is_err()
{
return false;
}
if let Ok(mut guard) = LOG_MANAGER.lock() {
guard.log_dir = Some(dir.clone());
guard.active_file = Some(active_file);
guard.lines.clear();
}
cleanup_old_logs(&dir);
true
}
#[napi]
pub fn configure_log_manager(core_log: bool, debug_log: bool) {
configure(core_log, debug_log);
}
#[napi]
pub fn write_app_log(level: i32, target: String, message: String) {
record_app_log(level, &target, &message);
}
#[napi]
pub fn drain_log_lines() -> Vec<String> {
LOG_MANAGER
.lock()
.map(|mut guard| guard.lines.drain(..).collect())
.unwrap_or_default()
}
#[napi]
pub fn list_log_files() -> Vec<LogFileInfo> {
let Some((log_dir, active_file)) = current_log_state() else {
return Vec::new();
};
let active_name = active_file.as_ref().and_then(|path| file_name(path));
let mut files = sorted_log_files(&log_dir);
files.reverse();
files
.into_iter()
.filter_map(|path| {
let file_name = file_name(&path)?;
let active = active_name
.as_ref()
.map(|name| name == &file_name)
.unwrap_or(false);
let metadata = fs::metadata(&path).ok();
Some(LogFileInfo {
file_name,
display_name: if active {
"当前启动日志".to_string()
} else {
"历史日志".to_string()
},
size_bytes: metadata
.as_ref()
.map(|value| value.len().min(i64::MAX as u64) as i64)
.unwrap_or(0),
modified_ms: metadata.as_ref().map(modified_millis).unwrap_or_default(),
active,
})
})
.collect()
}
#[napi]
pub fn read_log_file(file_name: String) -> Option<String> {
let (log_dir, _) = current_log_state()?;
let path = resolve_log_file(&log_dir, &file_name)?;
fs::read_to_string(path).ok()
}
#[napi]
pub fn export_log_file(file_name: String, target_path: String) -> bool {
let Some((log_dir, _)) = current_log_state() else {
return false;
};
let Some(path) = resolve_log_file(&log_dir, &file_name) else {
return false;
};
fs::copy(path, target_path).is_ok()
}
#[napi]
pub fn export_log_archive(target_path: String) -> bool {
let log_dir = LOG_MANAGER
.lock()
.ok()
.and_then(|guard| guard.log_dir.clone());
let Some(log_dir) = log_dir else {
return false;
};
let files = sorted_log_files(&log_dir);
let mut output = match OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(&target_path)
{
Ok(file) => file,
Err(_) => return false,
};
for path in files {
let name = path
.file_name()
.and_then(|value| value.to_str())
.unwrap_or("unknown.log");
let _ = writeln!(output, "===== {} =====", name);
if let Ok(content) = fs::read_to_string(&path) {
let _ = writeln!(output, "{}", content);
}
}
true
}
@@ -1 +1,2 @@
pub(crate) mod log_manager;
pub(crate) mod native_log;
@@ -1,7 +1,5 @@
use super::log_manager;
use napi_derive_ohos::napi;
use ohos_hilog_binding::{
LogOptions, hilog_debug, hilog_error, hilog_info, hilog_warn, set_global_options,
};
use std::collections::HashMap;
use std::panic;
use tracing::{Event, Subscriber};
@@ -10,8 +8,9 @@ use tracing_subscriber::layer::{Context, Layer};
use tracing_subscriber::prelude::*;
static INITIALIZED: std::sync::Once = std::sync::Once::new();
static TRACING_INITIALIZED: std::sync::Once = std::sync::Once::new();
fn panic_hook(info: &panic::PanicHookInfo) {
hilog_error!("RUST PANIC: {}", info);
log_manager::record_core_log(5, "RustPanic", &format!("{}", info));
}
#[napi]
@@ -23,45 +22,40 @@ pub fn init_panic_hook() {
#[napi]
pub fn hilog_global_options(domain: u32, tag: String) {
ohos_hilog_binding::forward_stdio_to_hilog();
set_global_options(LogOptions {
domain,
tag: Box::leak(tag.clone().into_boxed_str()),
})
let _ = domain;
let _ = tag;
}
#[napi]
pub fn init_tracing_subscriber() {
tracing_subscriber::registry()
.with(CallbackLayer {
callback: Box::new(tracing_callback),
})
.init();
TRACING_INITIALIZED.call_once(|| {
let _ = tracing_subscriber::registry()
.with(CallbackLayer {
callback: Box::new(tracing_callback),
})
.try_init();
});
}
fn tracing_callback(event: &Event, fields: HashMap<String, String>) {
let metadata = event.metadata();
#[cfg(target_env = "ohos")]
{
let loc = metadata.target().split("::").last().unwrap();
match *metadata.level() {
Level::TRACE => {
hilog_debug!("[{}] {:?}", loc, fields.values().collect::<Vec<_>>());
}
Level::DEBUG => {
hilog_debug!("[{}] {:?}", loc, fields.values().collect::<Vec<_>>());
}
Level::INFO => {
hilog_info!("[{}] {:?}", loc, fields.values().collect::<Vec<_>>());
}
Level::WARN => {
hilog_warn!("[{}] {:?}", loc, fields.values().collect::<Vec<_>>());
}
Level::ERROR => {
hilog_error!("[{}] {:?}", loc, fields.values().collect::<Vec<_>>());
}
}
let loc = metadata
.target()
.split("::")
.last()
.unwrap_or(metadata.target());
let level = match *metadata.level() {
Level::TRACE => 2,
Level::DEBUG => 3,
Level::INFO => 4,
Level::WARN => 6,
Level::ERROR => 5,
};
if !log_manager::core_log_enabled(level) {
return;
}
let values = fields.values().cloned().collect::<Vec<_>>().join(" ");
log_manager::record_core_log(level, &format!("Rust:{}", loc), &values);
}
struct CallbackLayer {
@@ -70,6 +64,16 @@ struct CallbackLayer {
impl<S: Subscriber> Layer<S> for CallbackLayer {
fn on_event(&self, event: &Event, _ctx: Context<S>) {
let level = match *event.metadata().level() {
Level::TRACE => 2,
Level::DEBUG => 3,
Level::INFO => 4,
Level::WARN => 6,
Level::ERROR => 5,
};
if !log_manager::core_log_enabled(level) {
return;
}
// 使用 fmt::format::FmtSpan 提取字段值
let mut fields = HashMap::new();
let mut visitor = FieldCollector(&mut fields);
@@ -3,6 +3,7 @@ use napi_derive_ohos::napi;
use serde::Serialize;
use std::collections::HashSet;
use std::sync::Mutex;
use url::Url;
static ATTACHED_TUN_INSTANCE_IDS: once_cell::sync::Lazy<Mutex<HashSet<String>>> =
once_cell::sync::Lazy::new(|| Mutex::new(HashSet::new()));
@@ -158,6 +159,136 @@ fn stringify_uuid(value: Option<common::Uuid>) -> Option<String> {
value.map(|v| v.to_string())
}
fn non_empty_string(value: Option<String>) -> Option<String> {
value.and_then(|raw| {
let trimmed = raw.trim();
if trimmed.is_empty() {
None
} else {
Some(trimmed.to_string())
}
})
}
fn config_virtual_ipv4_cidr(config: &api::manage::NetworkConfig) -> Option<String> {
non_empty_string(config.virtual_ipv4.clone())
.map(|ipv4| format!("{}/{}", ipv4, config.network_length.unwrap_or(24)))
}
fn config_endpoint_urls(config: &api::manage::NetworkConfig) -> Vec<String> {
let mut urls = Vec::new();
let mut seen = HashSet::new();
if let Some(url) = non_empty_string(config.public_server_url.clone())
&& seen.insert(url.clone())
{
urls.push(url);
}
for raw in &config.peer_urls {
let trimmed = raw.trim();
if trimmed.is_empty() {
continue;
}
let value = trimmed.to_string();
if seen.insert(value.clone()) {
urls.push(value);
}
}
urls
}
fn endpoint_url(url: &str) -> Option<Url> {
Url::parse(url).ok()
}
fn endpoint_scheme(url: &str) -> Option<String> {
endpoint_url(url)
.map(|parsed| parsed.scheme().to_string())
.or_else(|| {
let scheme = url.split("://").next().unwrap_or("").trim();
(!scheme.is_empty()).then_some(scheme.to_string())
})
}
fn endpoint_label(url: &str) -> String {
if let Some(parsed) = endpoint_url(url)
&& let Some(host) = parsed.host_str()
{
return format!("[Config] {}", host);
}
format!("[Config] {}", url)
}
fn endpoint_remote_display(url: &str) -> String {
if let Some(parsed) = endpoint_url(url)
&& let Some(host) = parsed.host_str()
{
return parsed
.port()
.map(|port| format!("{}:{}", host, port))
.unwrap_or_else(|| host.to_string());
}
url.to_string()
}
fn configured_peer_id(index: usize) -> i64 {
9_000_000 + index as i64
}
fn configured_route_views(endpoints: &[String], public_server_url: Option<&str>) -> Vec<RouteView> {
endpoints
.iter()
.enumerate()
.map(|(index, endpoint)| RouteView {
peer_id: configured_peer_id(index),
hostname: Some(endpoint_label(endpoint)),
ipv4: Some(endpoint_remote_display(endpoint)),
ipv4_cidr: None,
ipv6_cidr: None,
proxy_cidrs: Vec::new(),
next_hop_peer_id: None,
cost: Some(0),
path_latency: None,
udp_nat_type: None,
tcp_nat_type: None,
inst_id: None,
version: None,
is_public_server: public_server_url.map(|url| url == endpoint),
})
.collect()
}
fn configured_peer_views(endpoints: &[String]) -> Vec<PeerInfo> {
endpoints
.iter()
.enumerate()
.map(|(index, endpoint)| {
let conn_id = format!("configured-peer-{}", index);
PeerInfo {
peer_id: configured_peer_id(index),
default_conn_id: Some(conn_id.clone()),
directly_connected_conns: vec![conn_id.clone()],
conns: vec![PeerConnInfo {
conn_id,
my_peer_id: 0,
peer_id: configured_peer_id(index),
features: Vec::new(),
tunnel_type: endpoint_scheme(endpoint),
local_addr: None,
remote_addr: Some(endpoint.clone()),
resolved_remote_addr: Some(endpoint_remote_display(endpoint)),
stats: None,
loss_rate: None,
is_client: true,
network_name: None,
is_closed: false,
secure_auth_level: None,
peer_identity_type: None,
}],
}
})
.collect()
}
fn optional_u32_to_i64(value: Option<u32>) -> Option<i64> {
value.map(|v| v as i64)
}
@@ -193,7 +324,7 @@ fn route_to_view(route: api::instance::Route) -> RouteView {
}
}
fn peer_conn_to_view(conn: api::instance::PeerConnInfo) -> PeerConnInfo {
pub(crate) 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,
@@ -291,3 +422,43 @@ pub fn runtime_instance_from_running_info(
peers: info.peers.into_iter().map(peer_to_view).collect(),
}
}
pub fn runtime_instance_from_config_snapshot(
config_id: String,
display_name: String,
config: api::manage::NetworkConfig,
running: bool,
) -> RuntimeInstanceState {
let tun_attached = running && is_tun_attached(&config_id);
let tun_required =
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 my_node_info = MyNodeInfo {
virtual_ipv4: non_empty_string(config.virtual_ipv4.clone()),
virtual_ipv4_cidr: config_virtual_ipv4_cidr(&config),
hostname: non_empty_string(config.hostname.clone()),
version: None,
peer_id: None,
listeners: config.listener_urls.clone(),
vpn_portal_cfg: None,
udp_nat_type: None,
tcp_nat_type: None,
};
RuntimeInstanceState {
config_id: config_id.clone(),
instance_id: config_id,
display_name,
running,
tun_required,
tun_attached,
magic_dns_enabled: config.enable_magic_dns.unwrap_or(false),
need_exit_node: !config.exit_nodes.is_empty(),
error_message: None,
my_node_info: Some(my_node_info),
events: Vec::new(),
routes: configured_route_views(&endpoint_urls, public_server_url.as_deref()),
peers: configured_peer_views(&endpoint_urls),
}
}
+43 -23
View File
@@ -654,7 +654,8 @@ mod manager {
#[derive(Default)]
pub(super) enum PersistedConfigSource {
User,
Webhook,
#[serde(alias = "webhook")]
Web,
#[serde(other)]
#[default]
Legacy,
@@ -664,15 +665,15 @@ mod manager {
pub(super) fn from_runtime_source(source: ConfigSource) -> Self {
match source {
ConfigSource::User => Self::User,
ConfigSource::Webhook => Self::Webhook,
ConfigSource::Web => Self::Web,
}
}
fn merge_persisted(self, incoming: Self) -> Self {
match (self, incoming) {
// Older runtimes report missing source as `user`. Keep the stronger persisted
// ownership until webhook sync or an explicit user save repairs it.
(Self::Webhook, Self::User) | (Self::Legacy, Self::User) => self,
// ownership until web sync or an explicit user save repairs it.
(Self::Web, Self::User) | (Self::Legacy, Self::User) => self,
(_, next) => next,
}
}
@@ -680,13 +681,13 @@ mod manager {
fn to_runtime_source(self) -> ConfigSource {
match self {
Self::User | Self::Legacy => ConfigSource::User,
Self::Webhook => ConfigSource::Webhook,
Self::Web => ConfigSource::Web,
}
}
#[cfg(any(test, target_os = "android"))]
fn is_webhook_like(self) -> bool {
matches!(self, Self::Webhook)
fn is_web_like(self) -> bool {
matches!(self, Self::Web)
}
}
@@ -918,7 +919,7 @@ mod manager {
}
#[cfg(target_os = "android")]
pub fn get_enabled_instances_with_webhook_like_tun_ids(
pub fn get_enabled_instances_with_web_like_tun_ids(
&self,
) -> impl Iterator<Item = uuid::Uuid> + '_ {
self.storage
@@ -926,7 +927,7 @@ mod manager {
.iter()
.filter(|v| self.storage.enabled_networks.contains(v.key()))
.filter(|v| !v.config.no_tun())
.filter(|v| v.source.is_webhook_like())
.filter(|v| v.source.is_web_like())
.filter_map(|c| c.config.instance_id().parse::<uuid::Uuid>().ok())
}
@@ -934,12 +935,11 @@ mod manager {
pub(super) async fn disable_instances_with_tun(
&self,
app: &AppHandle,
webhook_only: bool,
web_only: bool,
) -> Result<(), easytier::rpc_service::remote_client::RemoteClientError<anyhow::Error>>
{
let inst_ids: Vec<uuid::Uuid> = if webhook_only {
self.get_enabled_instances_with_webhook_like_tun_ids()
.collect()
let inst_ids: Vec<uuid::Uuid> = if web_only {
self.get_enabled_instances_with_web_like_tun_ids().collect()
} else {
self.get_enabled_instances_with_tun_ids().collect()
};
@@ -977,7 +977,7 @@ mod manager {
.await
.map_err(|e| e.to_string())?;
}
PersistedConfigSource::Webhook => {
PersistedConfigSource::Web => {
self.disable_instances_with_tun(app, true)
.await
.map_err(|e| e.to_string())?;
@@ -1187,26 +1187,46 @@ mod manager {
}
#[test]
fn persisted_source_merge_keeps_legacy_and_webhook_over_ambiguous_user() {
fn stored_gui_config_deserializes_webhook_source_as_web() {
let stored: StoredGuiConfig = serde_json::from_value(serde_json::json!({
"config": NetworkConfig::default(),
"source": "webhook",
}))
.unwrap();
assert_eq!(stored.source, PersistedConfigSource::Web);
}
#[test]
fn stored_gui_config_defaults_unknown_source_to_legacy() {
let stored: StoredGuiConfig = serde_json::from_value(serde_json::json!({
"config": NetworkConfig::default(),
"source": "unknown",
}))
.unwrap();
assert_eq!(stored.source, PersistedConfigSource::Legacy);
}
#[test]
fn persisted_source_merge_keeps_legacy_and_web_over_ambiguous_user() {
assert_eq!(
PersistedConfigSource::Legacy.merge_persisted(PersistedConfigSource::User),
PersistedConfigSource::Legacy
);
assert_eq!(
PersistedConfigSource::Webhook.merge_persisted(PersistedConfigSource::User),
PersistedConfigSource::Webhook
PersistedConfigSource::Web.merge_persisted(PersistedConfigSource::User),
PersistedConfigSource::Web
);
assert_eq!(
PersistedConfigSource::Legacy.merge_persisted(PersistedConfigSource::Webhook),
PersistedConfigSource::Webhook
PersistedConfigSource::Legacy.merge_persisted(PersistedConfigSource::Web),
PersistedConfigSource::Web
);
}
#[test]
fn only_webhook_configs_are_webhook_like() {
assert!(!PersistedConfigSource::Legacy.is_webhook_like());
assert!(!PersistedConfigSource::User.is_webhook_like());
assert!(PersistedConfigSource::Webhook.is_webhook_like());
fn only_web_configs_are_web_like() {
assert!(!PersistedConfigSource::Legacy.is_web_like());
assert!(!PersistedConfigSource::User.is_web_like());
assert!(PersistedConfigSource::Web.is_web_like());
}
}
}
+3 -4
View File
@@ -1,12 +1,11 @@
import { invoke } from '@tauri-apps/api/core'
import { Api, NetworkTypes } from 'easytier-frontend-lib'
import { GetNetworkMetasResponse } from 'node_modules/easytier-frontend-lib/dist/modules/api'
import { type ConfigSource, normalizeConfigSource } from './config_source'
type NetworkConfig = NetworkTypes.NetworkConfig
type ValidateConfigResponse = Api.ValidateConfigResponse
type ListNetworkInstanceIdResponse = Api.ListNetworkInstanceIdResponse
type ConfigSource = 'user' | 'webhook' | 'legacy'
interface ServiceOptions {
config_dir: string
rpc_portal: string
@@ -32,14 +31,14 @@ function parseStoredConfigs(raw: string | null): StoredGuiConfig[] {
if (entry && typeof entry === 'object' && 'config' in entry) {
const { config, source } = entry as {
config?: NetworkConfig
source?: ConfigSource
source?: unknown
}
if (!config) {
return []
}
return [{
config: NetworkTypes.normalizeNetworkConfig(config),
source: source === 'user' || source === 'webhook' ? source : 'legacy',
source: normalizeConfigSource(source),
}]
}
@@ -0,0 +1,13 @@
export type ConfigSource = 'user' | 'web' | 'legacy'
export function normalizeConfigSource(source: unknown): ConfigSource {
if (source === 'user' || source === 'web' || source === 'legacy') {
return source
}
if (source === 'webhook') {
return 'web'
}
return 'legacy'
}
+3 -2
View File
@@ -2,10 +2,11 @@ import { Event, listen } from "@tauri-apps/api/event";
import { type } from "@tauri-apps/plugin-os";
import { NetworkTypes } from "easytier-frontend-lib"
import { Utils } from "easytier-frontend-lib";
import { normalizeConfigSource } from './config_source'
interface StoredGuiConfig {
config: NetworkTypes.NetworkConfig
source?: 'user' | 'webhook' | 'legacy'
source?: unknown
}
const EVENTS = Object.freeze({
@@ -24,7 +25,7 @@ function onSaveConfigs(event: Event<StoredGuiConfig[]>) {
'networkList',
JSON.stringify(event.payload.map(({ config, source }) => ({
config: NetworkTypes.normalizeNetworkConfig(config),
source: source ?? 'legacy',
source: normalizeConfigSource(source),
}))),
);
}
+18 -15
View File
@@ -20,7 +20,7 @@ use session::{Location, Session};
use storage::{Storage, StorageToken};
use crate::FeatureFlags;
use crate::webhook::SharedWebhookConfig;
use crate::webhook::{ManagedNetworkConfig, SharedWebhookConfig};
use tokio::task::JoinSet;
use crate::db::{Db, UserIdInDb, entity::user_running_network_configs};
@@ -146,20 +146,7 @@ impl ClientManager {
}
pub async fn list_sessions(&self) -> Vec<StorageToken> {
let sessions = self
.client_sessions
.iter()
.map(|item| item.value().clone())
.collect::<Vec<_>>();
let mut ret: Vec<StorageToken> = vec![];
for s in sessions {
if let Some(t) = s.get_token().await {
ret.push(t);
}
}
ret
self.storage.list_clients()
}
pub fn get_session_by_machine_id(
@@ -197,6 +184,22 @@ impl ClientManager {
self.storage.list_user_clients(user_id)
}
pub async fn reconcile_managed_network_configs(
&self,
user_id: UserIdInDb,
machine_id: uuid::Uuid,
desired_configs: Vec<ManagedNetworkConfig>,
) -> anyhow::Result<()> {
session::SessionRpcService::reconcile_web_source_configs(
&self.storage,
user_id,
machine_id,
desired_configs,
)
.await?;
Ok(())
}
pub async fn get_heartbeat_requests(&self, client_url: &url::Url) -> Option<HeartbeatRequest> {
let s = self.client_sessions.get(client_url)?.clone();
s.data().read().await.req()
File diff suppressed because it is too large Load Diff
@@ -114,6 +114,20 @@ impl Storage {
.unwrap_or_default()
}
pub fn list_clients(&self) -> Vec<StorageToken> {
self.0
.user_clients_map
.iter()
.flat_map(|user_clients| {
user_clients
.value()
.iter()
.map(|info| info.value().storage_token.clone())
.collect::<Vec<_>>()
})
.collect()
}
pub fn db(&self) -> &Db {
&self.0.db
}
@@ -174,4 +188,25 @@ mod tests {
assert_eq!(storage.get_client_url_by_machine_id(2, &machine_id), None);
}
#[tokio::test]
async fn list_clients_returns_current_storage_tokens() {
let storage = Storage::new(Db::memory_db().await);
let user1_token = make_storage_token(1, uuid::Uuid::new_v4(), "tcp://127.0.0.1:1001");
let user2_token = make_storage_token(2, uuid::Uuid::new_v4(), "tcp://127.0.0.1:1002");
storage.update_client(user1_token.clone(), 10);
storage.update_client(user2_token.clone(), 20);
let tokens = storage.list_clients();
assert_eq!(tokens.len(), 2);
assert!(tokens.iter().any(|token| token.token == user1_token.token));
assert!(tokens.iter().any(|token| token.token == user2_token.token));
storage.remove_client(&user1_token);
let tokens = storage.list_clients();
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].token, user2_token.token);
}
}
+6 -6
View File
@@ -331,7 +331,7 @@ mod tests {
(user_id, device_id),
inst_id,
network_config,
ConfigSource::Webhook,
ConfigSource::Web,
)
.await
.unwrap();
@@ -344,10 +344,10 @@ mod tests {
.unwrap();
println!("device: {}, {:?}", device_id, result2);
assert_eq!(result2.network_config, network_config_json);
assert_eq!(result2.get_network_config_source(), ConfigSource::Webhook);
assert_eq!(result2.get_network_config_source(), ConfigSource::Web);
assert_eq!(
result2.get_runtime_network_config_source(),
ConfigSource::Webhook
ConfigSource::Web
);
assert_eq!(result.create_time, result2.create_time);
@@ -373,7 +373,7 @@ mod tests {
}
#[tokio::test]
async fn test_legacy_network_config_defaults_to_user_runtime_source() {
async fn test_unknown_network_config_source_defaults_to_user_runtime_source() {
let db = Db::memory_db().await;
let user_id = 1;
let inst_id = uuid::Uuid::new_v4();
@@ -384,11 +384,11 @@ mod tests {
device_id: Set(device_id.to_string()),
network_instance_id: Set(inst_id.to_string()),
network_config: Set(serde_json::to_string(&NetworkConfig {
network_name: Some("legacy".to_string()),
network_name: Some("unknown-source".to_string()),
..Default::default()
})
.unwrap()),
source: Set("legacy".to_string()),
source: Set("unknown".to_string()),
disabled: Set(false),
create_time: Set(sqlx::types::chrono::Local::now().fixed_offset()),
update_time: Set(sqlx::types::chrono::Local::now().fixed_offset()),
@@ -48,7 +48,7 @@ impl MigrationTrait for Migration {
device_id,
network_instance_id,
network_config,
'legacy',
'user',
disabled,
create_time,
update_time
@@ -0,0 +1,42 @@
use sea_orm_migration::prelude::*;
pub struct Migration;
impl MigrationName for Migration {
fn name(&self) -> &str {
"m20260514_000004_rename_web_config_source"
}
}
#[async_trait::async_trait]
impl MigrationTrait for Migration {
async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> {
let db = manager.get_connection();
db.execute_unprepared(
r#"
UPDATE user_running_network_configs
SET source = 'web'
WHERE source = 'webhook';
UPDATE user_running_network_configs
SET source = 'user'
WHERE source = 'legacy';
"#,
)
.await?;
Ok(())
}
async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> {
let db = manager.get_connection();
db.execute_unprepared(
r#"
UPDATE user_running_network_configs
SET source = 'webhook'
WHERE source = 'web';
"#,
)
.await?;
Ok(())
}
}
+2
View File
@@ -3,6 +3,7 @@ use sea_orm_migration::prelude::*;
mod m20241029_000001_init;
mod m20260403_000002_scope_network_config_unique;
mod m20260421_000003_add_network_config_source;
mod m20260514_000004_rename_web_config_source;
pub struct Migrator;
@@ -13,6 +14,7 @@ impl MigratorTrait for Migrator {
Box::new(m20241029_000001_init::Migration),
Box::new(m20260403_000002_scope_network_config_unique::Migration),
Box::new(m20260421_000003_add_network_config_source::Migration),
Box::new(m20260514_000004_rename_web_config_source::Migration),
]
}
}
+53 -3
View File
@@ -3,6 +3,7 @@ use axum::http::StatusCode;
use axum::routing::{delete, post};
use axum::{Json, Router, extract::State, routing::get};
use axum_login::AuthUser;
use easytier::common::config::ConfigSource as RuntimeConfigSource;
use easytier::launcher::NetworkConfig;
use easytier::proto::common::Void;
use easytier::proto::{api::manage::*, web::*};
@@ -60,6 +61,7 @@ struct SaveNetworkJsonReq {
struct RunNetworkJsonReq {
config: NetworkConfig,
save: bool,
source: Option<i32>,
}
#[derive(Debug, serde::Deserialize, serde::Serialize)]
@@ -82,6 +84,17 @@ struct RemoveNetworkJsonReq {
inst_ids: Vec<uuid::Uuid>,
}
#[derive(Debug, serde::Deserialize, serde::Serialize)]
struct ManagedNetworkConfigJson {
instance_id: uuid::Uuid,
network_config: serde_json::Value,
}
#[derive(Debug, serde::Deserialize, serde::Serialize)]
struct ReconcileManagedNetworkConfigsJsonReq {
managed_network_configs: Vec<ManagedNetworkConfigJson>,
}
#[derive(Debug, serde::Deserialize, serde::Serialize)]
struct ListMachineItem {
client_url: Option<url::Url>,
@@ -130,10 +143,11 @@ impl NetworkApi {
Json(payload): Json<RunNetworkJsonReq>,
) -> Result<Json<Void>, HttpHandleError> {
client_mgr
.handle_run_network_instance(
.handle_run_network_instance_with_source(
(Self::get_user_id(&auth_session)?, machine_id),
payload.config,
payload.save,
RuntimeConfigSource::Web,
)
.await
.map_err(convert_error)?;
@@ -274,10 +288,11 @@ impl NetworkApi {
));
}
client_mgr
.handle_save_network_config(
.handle_save_network_config_with_source(
(Self::get_user_id(&auth_session)?, machine_id),
inst_id,
payload.config,
RuntimeConfigSource::Web,
)
.await
.map_err(convert_error)
@@ -302,8 +317,17 @@ impl NetworkApi {
Path((user_id, machine_id)): Path<(UserIdInDb, uuid::Uuid)>,
Json(payload): Json<RunNetworkJsonReq>,
) -> Result<Json<Void>, HttpHandleError> {
let source = payload
.source
.and_then(RuntimeConfigSource::from_rpc)
.unwrap_or(RuntimeConfigSource::Web);
client_mgr
.handle_run_network_instance((user_id, machine_id), payload.config, payload.save)
.handle_run_network_instance_with_source(
(user_id, machine_id),
payload.config,
payload.save,
source,
)
.await
.map_err(convert_error)?;
Ok(Void::default().into())
@@ -319,6 +343,31 @@ impl NetworkApi {
.map_err(convert_error)
}
async fn handle_reconcile_managed_network_configs_internal(
State(client_mgr): AppState,
Path((user_id, machine_id)): Path<(UserIdInDb, uuid::Uuid)>,
Json(payload): Json<ReconcileManagedNetworkConfigsJsonReq>,
) -> Result<Json<Void>, HttpHandleError> {
let desired = payload
.managed_network_configs
.into_iter()
.map(|item| crate::webhook::ManagedNetworkConfig {
instance_id: item.instance_id.to_string(),
network_config: item.network_config,
})
.collect();
client_mgr
.reconcile_managed_network_configs(user_id, machine_id, desired)
.await
.map_err(|err| {
(
StatusCode::INTERNAL_SERVER_ERROR,
other_error(err.to_string()).into(),
)
})?;
Ok(Void::default().into())
}
async fn handle_list_network_instance_ids_internal(
State(client_mgr): AppState,
Path((user_id, machine_id)): Path<(UserIdInDb, uuid::Uuid)>,
@@ -347,6 +396,7 @@ impl NetworkApi {
.route(
"/api/internal/users/:user-id/machines/:machine-id/networks",
post(Self::handle_run_network_instance_internal)
.put(Self::handle_reconcile_managed_network_configs_internal)
.get(Self::handle_list_network_instance_ids_internal),
)
.route(
+16 -6
View File
@@ -16,6 +16,7 @@ pub struct ProxyRpcRequest {
pub service_name: String,
pub method_name: String,
pub payload: serde_json::Value,
pub scope: Option<String>,
}
macro_rules! match_service {
@@ -35,6 +36,7 @@ async fn handle_proxy_rpc_by_session(
service_name,
method_name,
payload,
scope,
} = req;
let resp = match service_name.as_str() {
@@ -74,12 +76,20 @@ async fn handle_proxy_rpc_by_session(
payload,
session
),
"api.instance.TcpProxyRpcService" => match_service!(
easytier::proto::api::instance::TcpProxyRpcClientFactory<BaseController>,
method_name,
payload,
session
),
"api.instance.TcpProxyRpcService" => {
let client = if let Some(ref domain) = scope {
session.scoped_client_with_domain::<
easytier::proto::api::instance::TcpProxyRpcClientFactory<BaseController>,
>(domain.clone())
} else {
session.scoped_client::<
easytier::proto::api::instance::TcpProxyRpcClientFactory<BaseController>,
>()
};
client
.json_call_method(BaseController::default(), &method_name, payload)
.await
}
"api.instance.AclManageRpcService" => match_service!(
easytier::proto::api::instance::AclManageRpcClientFactory<BaseController>,
method_name,
+18 -1
View File
@@ -57,6 +57,8 @@ pub struct ValidateTokenRequest {
pub os_distribution: Option<String>,
pub web_instance_id: Option<String>,
pub web_instance_api_base_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub applied_config_revision: Option<String>,
}
#[derive(Debug, Deserialize)]
@@ -66,7 +68,8 @@ pub struct ValidateTokenResponse {
pub pre_approved: bool,
#[serde(default)]
pub binding_version: u64,
pub managed_network_configs: Vec<ManagedNetworkConfig>,
#[serde(default)]
pub managed_network_configs: Option<Vec<ManagedNetworkConfig>>,
pub config_revision: String,
}
@@ -184,3 +187,17 @@ impl WebhookConfig {
}
pub type SharedWebhookConfig = Arc<WebhookConfig>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_token_response_allows_missing_managed_configs() {
let resp: ValidateTokenResponse =
serde_json::from_str(r#"{"valid":true,"config_revision":"rev-1"}"#).unwrap();
assert!(resp.valid);
assert_eq!(resp.config_revision, "rev-1");
assert!(resp.managed_network_configs.is_none());
}
}
+22 -15
View File
@@ -51,7 +51,7 @@ time = "0.3"
toml = "0.8.12"
chrono = { version = "0.4.37", features = ["serde"] }
guarden = "0.1"
guarden = "0.2"
delegate = "0.13.5"
@@ -82,7 +82,8 @@ pin-project-lite = "0.2.13"
atomic_refcell = "0.1.13"
quinn = { version = "0.11.8", optional = true, features = ["ring"] }
quinn-plaintext = { version = "0.3.0", optional = true }
quinn-proto = { version = "0.11.12", optional = true }
seahash = { version = "4.1.0", optional = true }
rustls = { version = "0.23.0", features = [
"ring", "tls12"
@@ -90,7 +91,7 @@ rustls = { version = "0.23.0", features = [
rcgen = { version = "0.12.1", optional = true }
# for websocket
tokio-websockets = { version = "0.13.2", optional = true, features = [
tokio-websockets = { version = "0.13.2", git = "https://github.com/EasyTier/tokio-websockets", optional = true, features = [
"rustls-webpki-roots",
"client",
"server",
@@ -127,10 +128,13 @@ uuid = { version = "1.5.0", features = [
once_cell = "1.18.0"
# for rpc
prost = "0.13.5"
prost-wkt = "0.6"
prost-wkt-types = "0.6"
prost = "0.14.3"
prost-reflect = { version = "0.16.4", default-features = false, features = ["derive"] }
prost-wkt-types = "0.7.1"
pbjson = "0.9.0"
anyhow = "1.0"
ariadne = "0.5"
url = { version = "2.5", features = ["serde"] }
percent-encoding = "2.3.1"
@@ -211,6 +215,7 @@ smoltcp = { git = "https://github.com/smoltcp-rs/smoltcp.git", rev = "0a926767a6
"async",
] }
parking_lot = { version = "0.12.0" }
fastant = "0.1"
wildmatch = "2.3.4"
@@ -224,11 +229,7 @@ service-manager = { git = "https://github.com/EasyTier/service-manager-rs.git",
zstd = { version = "0.13", optional = true }
kcp-sys = { git = "https://github.com/EasyTier/kcp-sys", rev = "94964794caaed5d388463137da59b97499619e5f", optional = true }
prost-reflect = { version = "0.14.5", default-features = false, features = [
"derive",
] }
kcp-sys = { git = "https://github.com/EasyTier/kcp-sys", rev = "d7427c22d764deb1860a7d37acc446ed5033464c", optional = true }
# for http connector
http_req = { git = "https://github.com/EasyTier/http_req.git", default-features = false, features = [
@@ -320,9 +321,9 @@ cfg_aliases = "0.2.1"
indoc = "2.0"
globwalk = "0.8.1"
regex = "1"
prost-build = "0.13.5"
prost-wkt-build = "0.6"
prost-reflect-build = { version = "0.14.0" }
prost-build = "0.14.3"
prost-reflect-build = "0.16.0"
pbjson-build = "0.9.0"
proc-macro2 = "1"
quote = "1"
thunk-rs = { git = "https://github.com/easytier/thunk.git", default-features = false, features = [
@@ -341,6 +342,11 @@ futures-util = "0.3.31"
maplit = "1.0.2"
tempfile = "3.22.0"
ctor = "0.8.0"
criterion = { version = "0.5", features = ["html_reports"] }
[[bench]]
name = "counter_contention"
harness = false
[target.'cfg(target_os = "linux")'.dev-dependencies]
defguard_wireguard_rs = "0.4.2"
@@ -375,7 +381,7 @@ full = [
"zstd",
]
wireguard = ["dep:boringtun", "dep:ring"]
quic = ["dep:quinn", "dep:quinn-plaintext", "dep:rustls", "dep:rcgen"]
quic = ["dep:quinn", "dep:quinn-proto", "dep:seahash", "dep:rustls", "dep:rcgen"]
kcp = ["dep:kcp-sys"]
mimalloc = ["dep:mimalloc"]
aes-gcm = ["dep:aes-gcm"]
@@ -391,6 +397,7 @@ websocket = [
]
smoltcp = ["dep:smoltcp"]
socks5 = ["smoltcp"]
ffi-dataplane = ["socks5"]
jemalloc = ["dep:jemallocator", "dep:jemalloc-sys"]
jemalloc-prof = [
"jemalloc",
+443
View File
@@ -0,0 +1,443 @@
//! Compare counter implementations under tokio-task contention.
//!
//! Groups:
//! - `contention_scaling` : N tokio tasks share one counter, total work fixed.
//! Variants: `single_atomic`, `cas_saturating`, `sharded_atomic`,
//! `thread_local_cell`, `unsafe_cell` (unsound, for reference only).
//! - `single_thread_write`: per-`add` cost with no contention (floor cost).
//! - `read_cost` : per-`get()` cost.
//! - `counter_handle` : the REAL production hot path. Measures the actual
//! `stats_manager::CounterHandle::add` (single-atomic `fetch_add` + lock-free
//! fastant `touch`) against reconstructed baselines:
//! * `prod` - real `CounterHandle` (this code's version)
//! * `baseline_cas_mutex` - pre-optimization: single `AtomicU64` with `fetch_update` (CAS) + `Mutex<Instant>` touch
//! * `baseline_fetchadd_mutex` - `fetch_add` + `Mutex<Instant>` touch (isolates the lock-free fastant touch)
//!
//! Run: `cargo bench -p easytier --bench counter_contention`
use std::cell::{Cell, UnsafeCell};
use std::hint::black_box;
use std::sync::Arc;
use std::sync::LazyLock;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::thread::available_parallelism;
use std::time::Instant;
use criterion::{
BenchmarkId, Criterion, Throughput, criterion_group, criterion_main, measurement::WallTime,
};
use easytier::common::stats_manager::{CounterHandle, MetricName, StatsManager};
use parking_lot::Mutex;
const COUNTER_SHARDS: usize = 16;
const TOTAL_WORK: u64 = 8_000_000;
// The handle path does a counter update plus a timestamp `touch` per `add`,
// so it is heavier per op than the counter-only groups; use a smaller total to
// keep the bench fast.
const HANDLE_TOTAL_WORK: u64 = 2_000_000;
const TASK_COUNTS: &[usize] = &[1, 2, 4, 8, 16, 32];
trait Counter: Send + Sync {
fn add(&self, delta: u64);
fn get(&self) -> u64;
}
// ---------------------------------------------------------------------------
// 1. SingleAtomic: one atomic, fetch_add. Baseline; contends across cores.
// ---------------------------------------------------------------------------
struct SingleAtomic(AtomicU64);
impl Default for SingleAtomic {
fn default() -> Self {
Self(AtomicU64::new(0))
}
}
impl Counter for SingleAtomic {
#[inline(always)]
fn add(&self, delta: u64) {
self.0.fetch_add(delta, Ordering::Relaxed);
}
#[inline(always)]
fn get(&self) -> u64 {
self.0.load(Ordering::Relaxed)
}
}
// ---------------------------------------------------------------------------
// 2. CasSaturating: fetch_update with saturating_add (the original PR `add`).
// A CAS loop that can retry under contention.
// ---------------------------------------------------------------------------
struct CasSaturating(AtomicU64);
impl Default for CasSaturating {
fn default() -> Self {
Self(AtomicU64::new(0))
}
}
impl Counter for CasSaturating {
#[inline(always)]
fn add(&self, delta: u64) {
let _ = self
.0
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |c| {
Some(c.saturating_add(delta))
});
}
#[inline(always)]
fn get(&self) -> u64 {
self.0.load(Ordering::Relaxed)
}
}
// ---------------------------------------------------------------------------
// 3. ShardedAtomic: 16 cache-aligned shards + per-thread shard index.
// Comparison-only variant (production `stats_manager::Counter` is
// single-atomic; sharding was evaluated and dropped as no benefit for the
// default 1-16 worker deployments).
// ---------------------------------------------------------------------------
thread_local! {
static SHARD_IDX: Cell<usize> = Cell::new({
static NEXT: AtomicUsize = AtomicUsize::new(0);
NEXT.fetch_add(1, Ordering::Relaxed) % COUNTER_SHARDS
});
}
#[repr(align(64))]
struct Shard {
value: AtomicU64,
}
struct ShardedAtomic {
shards: Box<[Shard]>,
}
impl Default for ShardedAtomic {
fn default() -> Self {
let mut shards = Vec::with_capacity(COUNTER_SHARDS);
for _ in 0..COUNTER_SHARDS {
shards.push(Shard {
value: AtomicU64::new(0),
});
}
Self {
shards: shards.into_boxed_slice(),
}
}
}
impl Counter for ShardedAtomic {
#[inline(always)]
fn add(&self, delta: u64) {
let i = SHARD_IDX.with(|c| c.get());
self.shards[i].value.fetch_add(delta, Ordering::Relaxed);
}
#[inline(always)]
fn get(&self) -> u64 {
self.shards
.iter()
.map(|s| s.value.load(Ordering::Relaxed))
.sum()
}
}
// ---------------------------------------------------------------------------
// 4. ThreadLocalCell: per-thread Cell<u64> accumulation. Zero-atomic writes.
// `get()` flushes the caller thread's local into a shared aggregate, so the
// measured read cost reflects a flush-based read. Exact totals would require
// flushing every thread (not modeled here).
// ---------------------------------------------------------------------------
thread_local! {
static TLS_DELTA: Cell<u64> = const { Cell::new(0) };
}
struct ThreadLocalCell {
shared: AtomicU64,
}
impl Default for ThreadLocalCell {
fn default() -> Self {
Self {
shared: AtomicU64::new(0),
}
}
}
impl Counter for ThreadLocalCell {
#[inline(always)]
fn add(&self, delta: u64) {
TLS_DELTA.with(|c| c.set(c.get() + delta));
}
#[inline(always)]
fn get(&self) -> u64 {
let local = TLS_DELTA.with(|c| c.replace(0));
self.shared.fetch_add(local, Ordering::Relaxed) + local
}
}
// ---------------------------------------------------------------------------
// 5. UnsafeCellCounter: a plain u64 mutated through UnsafeCell with manual
// `unsafe impl Send/Sync`. This is UNSOUND under concurrent access (data
// race / UB) and is exactly what the original code did "for speed". It is
// included only to measure the speed ceiling the author was chasing, and to
// show that its `get()` returns wrong totals under contention (lost updates).
// ---------------------------------------------------------------------------
struct UnsafeCellCounter(UnsafeCell<u64>);
// SAFETY: deliberately unsound; see above.
unsafe impl Send for UnsafeCellCounter {}
unsafe impl Sync for UnsafeCellCounter {}
impl Default for UnsafeCellCounter {
fn default() -> Self {
Self(UnsafeCell::new(0))
}
}
impl Counter for UnsafeCellCounter {
#[inline(always)]
fn add(&self, delta: u64) {
// SAFETY: UNSOUND under concurrent access (data race).
unsafe {
*self.0.get() += delta;
}
}
#[inline(always)]
fn get(&self) -> u64 {
// SAFETY: UNSOUND under concurrent writers (data race).
unsafe { *self.0.get() }
}
}
// ---------------------------------------------------------------------------
// Production counter handle + reconstructed baselines for the `counter_handle`
// group. These measure the full hot path (`add` = counter update + `touch`
// timestamp), which is what actually runs per packet in `peer_manager`.
// ---------------------------------------------------------------------------
// The real production `CounterHandle`. `CounterHandle::add` does a single
// `AtomicU64::fetch_add` then a lock-free fastant `touch`.
impl Counter for CounterHandle {
#[inline(always)]
fn add(&self, delta: u64) {
// Fully-qualified to avoid infinite recursion through the trait method.
CounterHandle::add(self, delta);
}
#[inline(always)]
fn get(&self) -> u64 {
CounterHandle::get(self)
}
}
// A faithful replica of the pre-optimization design: a SINGLE `AtomicU64`
// (unsharded) plus a `Mutex<Instant>` timestamp. `use_cas` selects whether the
// counter write is a `fetch_update` saturating CAS (the original PR `add`) or a
// plain `fetch_add`.
struct BaselineHandle {
counter: AtomicU64,
last_updated: Mutex<Instant>,
use_cas: bool,
}
impl Default for BaselineHandle {
fn default() -> Self {
Self {
counter: AtomicU64::new(0),
last_updated: Mutex::new(Instant::now()),
use_cas: false,
}
}
}
impl Counter for BaselineHandle {
#[inline(always)]
fn add(&self, delta: u64) {
if self.use_cas {
let _ = self
.counter
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |c| {
Some(c.saturating_add(delta))
});
} else {
self.counter.fetch_add(delta, Ordering::Relaxed);
}
*self.last_updated.lock() = Instant::now();
}
#[inline(always)]
fn get(&self) -> u64 {
self.counter.load(Ordering::Relaxed)
}
}
// ---------------------------------------------------------------------------
// Harness: a shared multi-thread tokio runtime sized to the host's parallelism.
// ---------------------------------------------------------------------------
static RUNTIME: LazyLock<tokio::runtime::Runtime> = LazyLock::new(|| {
let workers = available_parallelism().map(|n| n.get()).unwrap_or(1);
tokio::runtime::Builder::new_multi_thread()
.worker_threads(workers)
.enable_all()
.build()
.expect("failed to build tokio runtime")
});
fn bench_contention<C: Counter + Default + 'static>(
group: &mut criterion::BenchmarkGroup<'_, WallTime>,
name: &str,
n_tasks: usize,
per_task: u64,
) {
let counter: Arc<C> = Arc::new(C::default());
group.bench_with_input(BenchmarkId::new(name, n_tasks), &n_tasks, |b, &n| {
b.iter(|| {
let counter = counter.clone();
RUNTIME.block_on(async move {
let mut handles = Vec::with_capacity(n);
for _ in 0..n {
let c = counter.clone();
handles.push(tokio::spawn(async move {
for _ in 0..per_task {
c.add(black_box(1));
}
}));
}
for handle in handles {
let _ = handle.await;
}
black_box(counter.get());
});
});
});
}
fn contention_scaling(c: &mut Criterion) {
let mut group = c.benchmark_group("contention_scaling");
group.throughput(Throughput::Elements(TOTAL_WORK));
for &n in TASK_COUNTS {
let per = TOTAL_WORK / n as u64;
bench_contention::<SingleAtomic>(&mut group, "single_atomic", n, per);
bench_contention::<CasSaturating>(&mut group, "cas_saturating", n, per);
bench_contention::<ShardedAtomic>(&mut group, "sharded_atomic", n, per);
bench_contention::<ThreadLocalCell>(&mut group, "thread_local_cell", n, per);
bench_contention::<UnsafeCellCounter>(&mut group, "unsafe_cell", n, per);
}
group.finish();
}
fn single_thread_write<C: Counter + Default>(
group: &mut criterion::BenchmarkGroup<'_, WallTime>,
name: &str,
) {
let counter = C::default();
group.bench_function(name, |b| {
b.iter(|| {
counter.add(black_box(1));
});
});
}
fn single_thread_write_group(c: &mut Criterion) {
let mut group = c.benchmark_group("single_thread_write");
group.throughput(Throughput::Elements(1));
single_thread_write::<SingleAtomic>(&mut group, "single_atomic");
single_thread_write::<CasSaturating>(&mut group, "cas_saturating");
single_thread_write::<ShardedAtomic>(&mut group, "sharded_atomic");
single_thread_write::<ThreadLocalCell>(&mut group, "thread_local_cell");
single_thread_write::<UnsafeCellCounter>(&mut group, "unsafe_cell");
group.finish();
}
fn read_cost<C: Counter + Default>(
group: &mut criterion::BenchmarkGroup<'_, WallTime>,
name: &str,
) {
let counter = C::default();
counter.add(1000);
group.bench_function(name, |b| {
b.iter(|| black_box(counter.get()));
});
}
fn read_cost_group(c: &mut Criterion) {
let mut group = c.benchmark_group("read_cost");
read_cost::<SingleAtomic>(&mut group, "single_atomic");
read_cost::<CasSaturating>(&mut group, "cas_saturating");
read_cost::<ShardedAtomic>(&mut group, "sharded_atomic");
read_cost::<ThreadLocalCell>(&mut group, "thread_local_cell");
read_cost::<UnsafeCellCounter>(&mut group, "unsafe_cell");
group.finish();
}
fn bench_handle(
group: &mut criterion::BenchmarkGroup<'_, WallTime>,
name: &str,
n_tasks: usize,
per_task: u64,
counter: Arc<dyn Counter>,
) {
group.bench_with_input(BenchmarkId::new(name, n_tasks), &n_tasks, |b, &n| {
b.iter(|| {
let counter = counter.clone();
RUNTIME.block_on(async move {
let mut handles = Vec::with_capacity(n);
for _ in 0..n {
let c = counter.clone();
handles.push(tokio::spawn(async move {
for _ in 0..per_task {
c.add(black_box(1));
}
}));
}
for handle in handles {
let _ = handle.await;
}
black_box(counter.get());
});
});
});
}
fn counter_handle(c: &mut Criterion) {
let mut group = c.benchmark_group("counter_handle");
group.throughput(Throughput::Elements(HANDLE_TOTAL_WORK));
// StatsManager::new() spawns a background cleanup task, which needs a tokio
// runtime context; bind it to our shared RUNTIME for the lifetime of the
// group.
let _rt_guard = RUNTIME.enter();
let stats = StatsManager::new();
let prod: Arc<dyn Counter> = Arc::new(stats.get_simple_counter(MetricName::TrafficBytesTx));
let cas: Arc<dyn Counter> = Arc::new(BaselineHandle {
use_cas: true,
..Default::default()
});
let fam: Arc<dyn Counter> = Arc::new(BaselineHandle {
use_cas: false,
..Default::default()
});
for &n in TASK_COUNTS {
let per = HANDLE_TOTAL_WORK / n as u64;
bench_handle(&mut group, "prod", n, per, prod.clone());
bench_handle(&mut group, "baseline_cas_mutex", n, per, cas.clone());
bench_handle(&mut group, "baseline_fetchadd_mutex", n, per, fam.clone());
}
group.finish();
}
// Keep the default measurement config; pass CLI flags to speed up a run, e.g.
// `-- --measurement-time 2 --sample-size 30 --warm-up-time 500`.
criterion_group! {
name = benches;
config = Criterion::default();
targets = contention_scaling, single_thread_write_group, read_cost_group, counter_handle
}
criterion_main!(benches);
+9 -24
View File
@@ -2,7 +2,6 @@ mod rpc;
use crate::rpc::ServiceGenerator;
use cfg_aliases::cfg_aliases;
use prost_wkt_build::{FileDescriptorSet, Message as _};
#[cfg(target_os = "windows")]
use std::io::Cursor;
use std::{env, path::PathBuf};
@@ -174,32 +173,15 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("cargo:rerun-if-changed={proto_file}");
}
let out = PathBuf::from(env::var("OUT_DIR").unwrap());
let descriptor_file = out.join("descriptors.bin");
let out = PathBuf::from(env::var("OUT_DIR")?);
let descriptor = out.join("descriptors.bin");
let mut config = prost_build::Config::new();
config
.type_attribute(".", "#[derive(serde::Serialize,serde::Deserialize)]")
.extern_path(".google.protobuf.Any", "::prost_wkt_types::Any")
.extern_path(".google.protobuf.Timestamp", "::prost_wkt_types::Timestamp")
.extern_path(".google.protobuf.Value", "::prost_wkt_types::Value")
.file_descriptor_set_path(&descriptor_file)
.protoc_arg("--experimental_allow_proto3_optional")
.type_attribute("peer_rpc.DirectConnectedPeerInfo", "#[derive(Hash)]")
.type_attribute("peer_rpc.PeerInfoForGlobalMap", "#[derive(Hash)]")
.type_attribute("peer_rpc.ForeignNetworkRouteInfoKey", "#[derive(Hash, Eq)]")
.type_attribute(
"peer_rpc.RouteForeignNetworkSummary.Info",
"#[derive(Hash, Eq)]",
)
.type_attribute("peer_rpc.RouteForeignNetworkSummary", "#[derive(Hash, Eq)]")
.type_attribute("common.RpcDescriptor", "#[derive(Hash, Eq)]")
.type_attribute("acl.Acl", "#[serde(default)]")
.type_attribute("acl.AclV1", "#[serde(default)]")
.type_attribute("acl.Chain", "#[serde(default)]")
.type_attribute("acl.Rule", "#[serde(default)]")
.type_attribute("acl.GroupInfo", "#[serde(default)]")
.field_attribute(".api.manage.NetworkConfig", "#[serde(default)]")
.file_descriptor_set_path(&descriptor)
.service_generator(Box::new(ServiceGenerator::default()))
.btree_map(["."])
.skip_debug([".common.Ipv4Addr", ".common.Ipv6Addr", ".common.UUID"]);
@@ -210,9 +192,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
.file_descriptor_set_bytes("crate::proto::DESCRIPTOR_POOL_BYTES")
.compile_protos_with_config(config, &proto_files_reflect, &["src/proto/"])?;
let descriptor_bytes = std::fs::read(descriptor_file).unwrap();
let descriptor = FileDescriptorSet::decode(&descriptor_bytes[..]).unwrap();
prost_wkt_build::add_serde(out, descriptor);
let descriptor = std::fs::read(descriptor)?;
pbjson_build::Builder::new()
.register_descriptors(&descriptor)?
.preserve_proto_field_names()
.btree_map(["."])
.build(&["."])?;
check_locale();
Ok(())
+7 -6
View File
@@ -5,12 +5,10 @@ core_clap:
en: |+
config server address, allow format:
full url: --config-server udp://127.0.0.1:22020/admin, 'udp' can be replaced with tcp, ws, wss (when config server ws is proxied to wss)
short link: --config-server https://example.com/easytier/admin, the HTTP(S) response should redirect to the full config server URL
only user name: --config-server admin, will use official server
zh-CN: |+
配置服务器地址。允许格式:
完整URL--config-server udp://127.0.0.1:22020/adminudp可以根据配置服务器替换为 tcp,ws,wss(配置服务器ws被代理为wss时)
短链接:--config-server https://example.com/easytier/adminHTTP(S) 响应应重定向到完整配置服务器 URL
仅用户名:--config-server admin,将使用官方的服务器
machine_id:
en: |+
@@ -116,11 +114,11 @@ core_clap:
en: "encryption algorithm to use, supported: '', 'xor', 'chacha20', 'aes-gcm', 'aes-gcm-256', 'openssl-aes128-gcm', 'openssl-aes256-gcm', 'openssl-chacha20'. Empty string means default (aes-gcm)"
zh-CN: "要使用的加密算法,支持:''(默认aes-gcm)、'xor'、'chacha20'、'aes-gcm'、'aes-gcm-256'、'openssl-aes128-gcm'、'openssl-aes256-gcm'、'openssl-chacha20'"
multi_thread:
en: "use multi-thread runtime, default is single-thread"
zh-CN: "使用多线程运行时默认为单线程"
en: "multi-thread tokio runtime (default on). Only affects launcher-based deployments (GUI/mobile/web/Windows service); the easytier-core CLI always runs single-threaded."
zh-CN: "多线程 tokio 运行时默认开启)。仅对 launcher 部署(GUI/移动端/web/Windows 服务)生效;easytier-core CLI 始终为单线程"
multi_thread_count:
en: "the number of threads to use, default is 2, only effective when multi-thread is enabled, must be greater than 2"
zh-CN: "使用的线程数,默认2,仅在多线程模式下有效。取值必须大于2"
en: "the number of worker threads, default 2, only effective when multi-thread is enabled, minimum 2"
zh-CN: "worker 线程数,默认 2,仅在启用多线程时生效,最小为 2"
disable_ipv6:
en: "do not use ipv6"
zh-CN: "不使用IPv6"
@@ -207,6 +205,9 @@ core_clap:
bind_device:
en: "bind the connector socket to physical devices to avoid routing issues. e.g.: subnet proxy segment conflicts with a node's segment, after binding the physical device, it can communicate with the node normally."
zh-CN: "将连接器的套接字绑定到物理设备以避免路由问题。比如子网代理网段与某节点的网段冲突,绑定物理设备后可以与该节点正常通信。"
socket_mark:
en: "Linux only: set SO_MARK (fwmark) on EasyTier's underlay sockets (TCP, UDP, QUIC, WebSocket, WireGuard, and the FakeTCP decoy socket) so the host can policy-route or filter them with 'ip rule fwmark ...', nftables ('meta mark'), or iptables ('-m mark'). Any value is applied verbatim (0 is a valid mark); omit the flag to leave SO_MARK untouched. Requires CAP_NET_ADMIN. Note: FakeTCP payload travels via raw TUN writes which the kernel does not tag — mark those separately on the TUN device if needed."
zh-CN: "仅 Linux: 在 EasyTier 的底层套接字 (TCP、UDP、QUIC、WebSocket、WireGuard 以及 FakeTCP 诱饵套接字) 上设置 SO_MARK (fwmark),使主机能用 'ip rule fwmark ...'、nftables ('meta mark') 或 iptables ('-m mark') 策略路由/过滤这些数据包。任何值都会原样应用 (0 也是合法的 mark);不传该参数即保持 SO_MARK 不变。需要 CAP_NET_ADMIN 权限。注意:FakeTCP 的实际载荷通过原始 TUN 写入,内核不会为其打标记;如有需要请在 TUN 设备上单独打标记。"
enable_kcp_proxy:
en: "proxy tcp streams with kcp, improving the latency and throughput on the network with udp packet loss."
zh-CN: "使用 KCP 代理 TCP 流,提高在 UDP 丢包网络上的延迟和吞吐量。"
+283 -43
View File
@@ -6,6 +6,7 @@ use std::{
};
use anyhow::Context;
use ariadne::{CharSet, Config as AriadneConfig, IndexType, Label, Report, ReportKind, Source};
use base64::{Engine as _, prelude::BASE64_STANDARD};
use clap::ValueEnum;
use clap::builder::PossibleValue;
@@ -73,6 +74,7 @@ pub fn gen_default_flags() -> Flags {
disable_upnp: false,
disable_relay_data: false,
enable_udp_broadcast_relay: false,
socket_mark: None,
}
}
@@ -277,20 +279,20 @@ pub struct NetworkIdentity {
pub enum ConfigSource {
#[default]
User,
Webhook,
Web,
}
impl ConfigSource {
pub fn as_str(self) -> &'static str {
match self {
Self::User => "user",
Self::Webhook => "webhook",
Self::Web => "web",
}
}
pub fn from_rpc(source: i32) -> Option<Self> {
match RpcConfigSource::try_from(source).ok() {
Some(RpcConfigSource::Webhook) => Some(Self::Webhook),
Some(RpcConfigSource::Web) => Some(Self::Web),
Some(RpcConfigSource::User) => Some(Self::User),
_ => None,
}
@@ -299,7 +301,7 @@ impl ConfigSource {
pub fn to_rpc(self) -> i32 {
match self {
Self::User => RpcConfigSource::User as i32,
Self::Webhook => RpcConfigSource::Webhook as i32,
Self::Web => RpcConfigSource::Web as i32,
}
}
}
@@ -310,7 +312,7 @@ impl std::str::FromStr for ConfigSource {
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"user" => Ok(Self::User),
"webhook" => Ok(Self::Webhook),
"web" => Ok(Self::Web),
other => Err(format!("unknown network config source: {other}")),
}
}
@@ -568,6 +570,35 @@ struct Config {
source: Option<ConfigSourceConfig>,
}
fn format_toml_parse_error(source_name: &str, config_str: &str, error: &toml::de::Error) -> String {
let message = format!("failed to parse config TOML from {source_name}");
let Some(span) = error.span() else {
return format!("{message}\ndetail: {error}");
};
let mut output = Vec::new();
let report = Report::build(ReportKind::Error, (source_name, span.clone()))
.with_config(
AriadneConfig::default()
.with_color(false)
.with_char_set(CharSet::Ascii)
.with_index_type(IndexType::Byte),
)
.with_message(&message)
.with_label(Label::new((source_name, span)).with_message(error.message()))
.finish();
if report
.write((source_name, Source::from(config_str)), &mut output)
.is_ok()
{
String::from_utf8_lossy(&output).into_owned()
} else {
format!("{message}\ndetail: {error}")
}
}
#[derive(Debug, Clone)]
pub struct TomlConfigLoader {
config: Arc<Mutex<Config>>,
@@ -590,50 +621,79 @@ impl TomlConfigLoader {
}
pub fn new_from_str(config_str: &str) -> Result<Self, anyhow::Error> {
let mut config = toml::de::from_str::<Config>(config_str)
.with_context(|| format!("failed to parse config file: {}", config_str))?;
Self::new_from_str_with_source("inline config", config_str)
}
pub fn new(config_path: &PathBuf) -> Result<Self, anyhow::Error> {
let config_str = std::fs::read_to_string(config_path)
.with_context(|| format!("failed to read config file: {}", config_path.display()))?;
let source_name = config_path.display().to_string();
Self::new_from_str_with_source(&source_name, &config_str)
}
pub(crate) fn new_from_str_with_source(
source_name: &str,
config_str: &str,
) -> Result<Self, anyhow::Error> {
let mut config = toml::de::from_str::<Config>(config_str).map_err(|err| {
let message = format_toml_parse_error(source_name, config_str, &err);
anyhow::Error::new(err).context(message)
})?;
Self::normalize_config_source(&mut config);
config.flags_struct = Some(Self::gen_flags(config.flags.clone().unwrap_or_default()));
Self::new_from_config(config).map_err(|err| {
let message = format!("failed to load config from {source_name}: {err}");
err.context(message)
})
}
fn new_from_config(mut config: Config) -> Result<Self, anyhow::Error> {
config.flags_struct = Some(
Self::gen_flags(config.flags.clone().unwrap_or_default())
.context("failed to parse flags")?,
);
let has_network_identity = config.network_identity.is_some();
let config = TomlConfigLoader {
config: Arc::new(Mutex::new(config)),
};
let old_ns = config.get_network_identity();
config.set_network_identity(NetworkIdentity::new(
old_ns.network_name,
old_ns.network_secret.unwrap_or_default(),
));
// Detect credential mode: secure_mode enabled + no network_secret in TOML
let is_credential = has_network_identity
&& config
.get_secure_mode()
.map(|sm| sm.enabled)
.unwrap_or(false)
&& old_ns
.network_secret
.as_deref()
.is_none_or(|s| s.is_empty());
if is_credential {
config.set_network_identity(NetworkIdentity::new_credential(old_ns.network_name));
} else {
config.set_network_identity(NetworkIdentity::new(
old_ns.network_name,
old_ns.network_secret.unwrap_or_default(),
));
}
Ok(config)
}
pub fn new(config_path: &PathBuf) -> Result<Self, anyhow::Error> {
let config_str = std::fs::read_to_string(config_path)
.with_context(|| format!("failed to read config file: {:?}", config_path))?;
let ret = Self::new_from_str(&config_str)?;
Ok(ret)
}
fn gen_flags(mut flags_hashmap: serde_json::Map<String, serde_json::Value>) -> Flags {
let default_flags_json = serde_json::to_string(&gen_default_flags()).unwrap();
let default_flags_hashmap =
serde_json::from_str::<serde_json::Map<String, serde_json::Value>>(&default_flags_json)
.unwrap();
let mut merged_hashmap = serde_json::Map::new();
for (key, value) in default_flags_hashmap {
if let Some(v) = flags_hashmap.remove(&key) {
merged_hashmap.insert(key, v);
} else {
merged_hashmap.insert(key, value);
}
}
serde_json::from_value(serde_json::Value::Object(merged_hashmap)).unwrap()
fn gen_flags(
flags_hashmap: serde_json::Map<String, serde_json::Value>,
) -> serde_json::Result<Flags> {
let mut merged_hashmap = match serde_json::to_value(gen_default_flags()) {
Ok(serde_json::Value::Object(map)) => map,
_ => serde_json::Map::new(),
};
merged_hashmap.extend(flags_hashmap);
serde_json::from_value(serde_json::Value::Object(merged_hashmap))
}
}
@@ -1190,13 +1250,13 @@ pub async fn load_config_from_file(
.read_to_string(&mut stdin)
.await
.context("failed to read config from stdin")?;
let config = TomlConfigLoader::new_from_str(&stdin)?;
let config = TomlConfigLoader::new_from_str_with_source("stdin", &stdin)?;
return Ok((config, ConfigFileControl::STATIC_CONFIG));
}
let config_str = tokio::fs::read_to_string(config_file)
.await
.with_context(|| format!("failed to read config file: {:?}", config_file))?;
.with_context(|| format!("failed to read config file: {}", config_file.display()))?;
let (expanded_config_str, uses_env_vars) = if disable_env_parsing {
(config_str.clone(), false)
@@ -1218,8 +1278,8 @@ pub async fn load_config_from_file(
);
}
let config = TomlConfigLoader::new_from_str(&expanded_config_str)
.with_context(|| format!("failed to load config file: {:?}", config_file))?;
let source_name = config_file.display().to_string();
let config = TomlConfigLoader::new_from_str_with_source(&source_name, &expanded_config_str)?;
let mut control = ConfigFileControl::from_path(config_file.clone()).await;
@@ -1259,6 +1319,147 @@ pub mod tests {
use std::path::PathBuf;
use tempfile::NamedTempFile;
#[test]
fn invalid_toml_error_includes_location_and_source_line() {
let error = TomlConfigLoader::new_from_str("dhcp = \"yes\"").unwrap_err();
let display = error.to_string();
assert!(display.contains("failed to parse config TOML"));
assert!(display.contains("inline config"));
assert!(display.contains("dhcp = \"yes\""));
assert!(display.contains("^"));
assert!(display.contains("invalid type: string"));
assert!(!display.contains("<unknown>"));
assert!(
error
.chain()
.any(|err| err.downcast_ref::<toml::de::Error>().is_some())
);
}
#[test]
fn invalid_file_toml_error_includes_config_source() {
let mut config_file = NamedTempFile::new().unwrap();
writeln!(config_file, "dhcp = \"yes\"").unwrap();
let error = TomlConfigLoader::new(&config_file.path().to_path_buf()).unwrap_err();
let error = error.to_string();
assert!(error.contains(config_file.path().to_string_lossy().as_ref()));
assert!(error.contains("failed to parse config TOML"));
assert!(error.contains("dhcp = \"yes\""));
assert!(error.contains("^"));
assert!(error.contains("invalid type: string"));
assert!(!error.contains("<unknown>"));
}
#[test]
fn invalid_stdin_toml_error_includes_config_source_in_display() {
let error = TomlConfigLoader::new_from_str_with_source("stdin", "dhcp = \"yes\"")
.unwrap_err()
.to_string();
assert!(error.contains("stdin"));
assert!(error.contains("failed to parse config TOML"));
assert!(error.contains("dhcp = \"yes\""));
assert!(error.contains("^"));
assert!(error.contains("invalid type: string"));
assert!(!error.contains("<unknown>"));
}
#[test]
fn invalid_toml_error_handles_non_ascii_before_error() {
let error = TomlConfigLoader::new_from_str("hostname = \"节点\"\ndhcp = \"yes\"")
.unwrap_err()
.to_string();
assert!(error.contains("dhcp = \"yes\""));
assert!(error.contains("^"));
assert!(error.contains("invalid type: string"));
}
#[test]
fn invalid_toml_error_handles_non_ascii_before_error_on_same_line() {
let error = TomlConfigLoader::new_from_str("hostname = \"节点\" dhcp = \"yes\"")
.unwrap_err()
.to_string();
assert!(error.contains("failed to parse config TOML"));
assert!(error.contains("inline config:1:"));
assert!(error.contains("hostname = \"节点\" dhcp = \"yes\""));
assert!(error.contains("expected newline"));
assert!(!error.contains("<unknown>"));
}
#[test]
fn invalid_file_flags_error_includes_config_source_in_display() {
let mut config_file = NamedTempFile::new().unwrap();
writeln!(config_file, "[flags]").unwrap();
writeln!(config_file, "socket_mark = \"bad\"").unwrap();
let error = TomlConfigLoader::new(&config_file.path().to_path_buf()).unwrap_err();
let display = error.to_string();
assert!(display.contains(config_file.path().to_string_lossy().as_ref()));
assert!(display.contains("failed to load config"));
assert!(display.contains("failed to parse flags"));
// with_context preserves the cause chain so callers can inspect the root reason.
let chain: Vec<String> = error.chain().map(|e| e.to_string()).collect();
assert!(chain.iter().any(|m| m.contains("failed to parse flags")));
}
#[test]
fn socket_mark_config_file_roundtrip_none_some_and_zero() {
// Omitting the flag leaves socket_mark unset (None) -> SO_MARK untouched.
let cfg = TomlConfigLoader::new_from_str(
r#"
[network_identity]
network_name = "n"
network_secret = "s"
"#,
)
.unwrap();
assert_eq!(cfg.get_flags().socket_mark, None);
// socket_mark = 0 is a legitimate value distinct from "unset".
let cfg = TomlConfigLoader::new_from_str(
r#"
[network_identity]
network_name = "n"
network_secret = "s"
[flags]
socket_mark = 0
"#,
)
.unwrap();
assert_eq!(cfg.get_flags().socket_mark, Some(0));
// A non-zero mark round-trips as Some(v).
let cfg = TomlConfigLoader::new_from_str(
r#"
[network_identity]
network_name = "n"
network_secret = "s"
[flags]
socket_mark = 66
"#,
)
.unwrap();
assert_eq!(cfg.get_flags().socket_mark, Some(66));
// set_flags(None) must serialize back through gen_config without
// resurrecting a value (guards the gen_flags merge against dropping
// the key when the serialized default is null).
cfg.set_flags(Flags {
socket_mark: None,
..cfg.get_flags()
});
assert_eq!(cfg.get_flags().socket_mark, None);
}
#[test]
fn test_stun_servers_config() {
let config = TomlConfigLoader::default();
@@ -1297,14 +1498,53 @@ stun_servers = [
let config = TomlConfigLoader::default();
assert_eq!(config.get_network_config_source(), ConfigSource::User);
config.set_network_config_source(Some(ConfigSource::Webhook));
config.set_network_config_source(Some(ConfigSource::Web));
let dumped = config.dump();
assert!(dumped.contains("[source]"));
assert!(dumped.contains("source = \"webhook\""));
assert!(dumped.contains("source = \"web\""));
let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap();
assert_eq!(loaded.get_network_config_source(), ConfigSource::Webhook);
assert_eq!(loaded.get_network_config_source(), ConfigSource::Web);
}
#[test]
fn test_toml_credential_mode_omits_network_secret() {
for network_secret in ["", r#"network_secret = """#] {
let config = TomlConfigLoader::new_from_str(&format!(
r#"
[network_identity]
network_name = "credential-network"
{network_secret}
[secure_mode]
enabled = true
"#
))
.unwrap();
let identity = config.get_network_identity();
assert_eq!(identity.network_name, "credential-network");
assert_eq!(identity.network_secret, None);
assert_eq!(identity.network_secret_digest, None);
assert!(!config.dump().contains("network_secret"));
}
}
#[test]
fn test_toml_secure_mode_without_network_identity_uses_default_secret() {
let config = TomlConfigLoader::new_from_str(
r#"
[secure_mode]
enabled = true
"#,
)
.unwrap();
let identity = config.get_network_identity();
assert_eq!(identity.network_name, "default");
assert_eq!(identity.network_secret.as_deref(), Some(""));
assert!(identity.network_secret_digest.is_some());
}
#[test]
+106 -115
View File
@@ -1,9 +1,10 @@
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use std::cell::UnsafeCell;
use std::fmt;
use std::sync::Arc;
use std::time::{Duration, Instant};
use std::sync::OnceLock;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use tokio::time::interval;
use tokio_util::task::AbortOnDropHandle;
@@ -374,136 +375,106 @@ impl Default for LabelSet {
}
}
/// UnsafeCounter provides a high-performance counter using UnsafeCell
/// Counter provides a high-performance atomic counter
#[derive(Debug)]
pub struct UnsafeCounter {
value: UnsafeCell<u64>,
pub struct Counter {
value: AtomicU64,
}
impl Default for UnsafeCounter {
impl Default for Counter {
fn default() -> Self {
Self::new()
}
}
impl UnsafeCounter {
impl Counter {
pub fn new() -> Self {
Self {
value: UnsafeCell::new(0),
value: AtomicU64::new(0),
}
}
pub fn new_with_value(initial: u64) -> Self {
Self {
value: UnsafeCell::new(initial),
value: AtomicU64::new(initial),
}
}
/// Increment the counter by the given amount
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn add(&self, delta: u64) {
let ptr = self.value.get();
unsafe {
*ptr = (*ptr).saturating_add(delta);
}
pub fn add(&self, delta: u64) {
self.value.fetch_add(delta, Ordering::Relaxed);
}
/// Increment the counter by 1
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn inc(&self) {
unsafe {
self.add(1);
}
pub fn inc(&self) {
self.add(1);
}
/// Get the current value of the counter
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is modifying this counter simultaneously.
pub unsafe fn get(&self) -> u64 {
let ptr = self.value.get();
unsafe { *ptr }
pub fn get(&self) -> u64 {
self.value.load(Ordering::Relaxed)
}
/// Reset the counter to zero
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn reset(&self) {
let ptr = self.value.get();
unsafe {
*ptr = 0;
}
pub fn reset(&self) {
self.value.store(0, Ordering::Relaxed);
}
/// Set the counter to a specific value
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this counter simultaneously.
pub unsafe fn set(&self, value: u64) {
let ptr = self.value.get();
unsafe {
*ptr = value;
}
pub fn set(&self, value: u64) {
self.value.store(value, Ordering::Relaxed);
}
}
// UnsafeCounter is Send + Sync because the safety is guaranteed by the caller
unsafe impl Send for UnsafeCounter {}
unsafe impl Sync for UnsafeCounter {}
/// Epoch used to convert a monotonic clock reading into a storable `u64`
/// millisecond count for `MetricData::last_updated`. Lazily initialized on first
/// use. Backed by `fastant`, which uses the TSC on x86_64 Linux (and falls back
/// to `std::time::Instant` elsewhere), making `now_millis()` cheap enough to
/// call per packet.
fn time_base() -> fastant::Instant {
static BASE: OnceLock<fastant::Instant> = OnceLock::new();
*BASE.get_or_init(fastant::Instant::now)
}
fn now_millis() -> u64 {
fastant::Instant::now()
.saturating_duration_since(time_base())
.as_millis() as u64
}
/// MetricData contains both the counter and last update timestamp
/// Uses UnsafeCell for lock-free access
#[derive(Debug)]
struct MetricData {
counter: UnsafeCounter,
last_updated: UnsafeCell<Instant>,
counter: Counter,
last_updated: AtomicU64,
}
impl MetricData {
fn new() -> Self {
Self {
counter: UnsafeCounter::new(),
last_updated: UnsafeCell::new(Instant::now()),
counter: Counter::new(),
last_updated: AtomicU64::new(now_millis()),
}
}
fn new_with_value(initial: u64) -> Self {
Self {
counter: UnsafeCounter::new_with_value(initial),
last_updated: UnsafeCell::new(Instant::now()),
counter: Counter::new_with_value(initial),
last_updated: AtomicU64::new(now_millis()),
}
}
/// Update the last_updated timestamp
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is accessing this timestamp simultaneously.
unsafe fn touch(&self) {
let ptr = self.last_updated.get();
unsafe {
*ptr = Instant::now();
}
/// Update the last_updated timestamp. Lock-free.
fn touch(&self) {
self.last_updated.store(now_millis(), Ordering::Relaxed);
}
/// Get the last updated timestamp
/// # Safety
/// This method is unsafe because it uses UnsafeCell. The caller must ensure
/// that no other thread is modifying this timestamp simultaneously.
unsafe fn get_last_updated(&self) -> Instant {
let ptr = self.last_updated.get();
unsafe { *ptr }
/// Last update time as milliseconds since `time_base()`.
fn last_updated_millis(&self) -> u64 {
self.last_updated.load(Ordering::Relaxed)
}
}
// MetricData is Send + Sync because the safety is guaranteed by the caller
unsafe impl Send for MetricData {}
unsafe impl Sync for MetricData {}
/// MetricKey uniquely identifies a metric with its name and labels
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
struct MetricKey {
@@ -546,39 +517,31 @@ impl CounterHandle {
/// Increment the counter by the given amount
pub fn add(&self, delta: u64) {
unsafe {
self.metric_data.counter.add(delta);
self.metric_data.touch();
}
self.metric_data.counter.add(delta);
self.metric_data.touch();
}
/// Increment the counter by 1
pub fn inc(&self) {
unsafe {
self.metric_data.counter.inc();
self.metric_data.touch();
}
self.metric_data.counter.inc();
self.metric_data.touch();
}
/// Get the current value of the counter
pub fn get(&self) -> u64 {
unsafe { self.metric_data.counter.get() }
self.metric_data.counter.get()
}
/// Reset the counter to zero
pub fn reset(&self) {
unsafe {
self.metric_data.counter.reset();
self.metric_data.touch();
}
self.metric_data.counter.reset();
self.metric_data.touch();
}
/// Set the counter to a specific value
pub fn set(&self, value: u64) {
unsafe {
self.metric_data.counter.set(value);
self.metric_data.touch();
}
self.metric_data.counter.set(value);
self.metric_data.touch();
}
}
@@ -614,9 +577,15 @@ impl StatsManager {
loop {
interval.tick().await;
let Some(cutoff_time) = Instant::now().checked_sub(Duration::from_secs(180)) else {
continue;
};
// Drop metrics untouched for 180s and with no live handles.
// Compare in the millis-since-base domain so neither the hot
// path nor GC reconstructs an `Instant` or locks.
//
// Use an age-based check (`now - last < STALE`) rather than
// `last > now - STALE`: early in process life `now_millis()` is
// tiny, so `now - STALE` saturates to 0 and a metric stamped at
// 0 would fail a strict `> 0` test and be wrongly evicted.
let now = now_millis();
let Some(counters) = counters_clone.upgrade() else {
break;
@@ -624,7 +593,7 @@ impl StatsManager {
counters.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|| now.saturating_sub(metric_data.last_updated_millis()) < 180_000
});
counters.shrink_to_fit();
}
@@ -662,7 +631,7 @@ impl StatsManager {
let key = entry.key();
let metric_data = entry.value();
let value = unsafe { metric_data.counter.get() };
let value = metric_data.counter.get();
metrics.push(MetricSnapshot {
name: key.name,
@@ -695,7 +664,7 @@ impl StatsManager {
let key = MetricKey::new(name, labels.clone());
if let Some(metric_data) = self.counters.get(&key) {
let value = unsafe { metric_data.counter.get() };
let value = metric_data.counter.get();
Some(MetricSnapshot {
name,
labels: labels.clone(),
@@ -793,20 +762,18 @@ mod tests {
}
#[tokio::test]
async fn test_unsafe_counter() {
let counter = UnsafeCounter::new();
async fn test_counter() {
let counter = Counter::new();
unsafe {
assert_eq!(counter.get(), 0);
counter.inc();
assert_eq!(counter.get(), 1);
counter.add(5);
assert_eq!(counter.get(), 6);
counter.set(10);
assert_eq!(counter.get(), 10);
counter.reset();
assert_eq!(counter.get(), 0);
}
assert_eq!(counter.get(), 0);
counter.inc();
assert_eq!(counter.get(), 1);
counter.add(5);
assert_eq!(counter.get(), 6);
counter.set(10);
assert_eq!(counter.get(), 10);
counter.reset();
assert_eq!(counter.get(), 0);
}
#[tokio::test]
@@ -947,12 +914,14 @@ mod tests {
let counter = stats.get_simple_counter(MetricName::TrafficBytesForwarded);
counter.set(1);
let cutoff_time = Instant::now().checked_add(Duration::from_secs(1)).unwrap();
// Cutoff 1s in the future, so every metric is stale by timestamp; only
// a live handle keeps a metric.
let cutoff_millis = now_millis() + 1_000;
stats
.counters
.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|| metric_data.last_updated_millis() > cutoff_millis
});
assert_eq!(stats.metric_count(), 1);
@@ -963,11 +932,33 @@ mod tests {
.counters
.retain(|_, metric_data: &mut Arc<MetricData>| {
Arc::strong_count(metric_data) > 1
|| unsafe { metric_data.get_last_updated() > cutoff_time }
|| metric_data.last_updated_millis() > cutoff_millis
});
assert_eq!(stats.metric_count(), 0);
}
#[tokio::test]
async fn test_counter_handle_concurrent_increment() {
const THREADS: usize = 8;
const INCREMENTS_PER_THREAD: usize = 10_000;
let stats = StatsManager::new();
let counter = stats.get_simple_counter(MetricName::TrafficPacketsForwarded);
std::thread::scope(|scope| {
for _ in 0..THREADS {
let counter = counter.clone();
scope.spawn(move || {
for _ in 0..INCREMENTS_PER_THREAD {
counter.inc();
}
});
}
});
assert_eq!(counter.get(), (THREADS * INCREMENTS_PER_THREAD) as u64);
}
#[tokio::test]
async fn test_stats_rpc_data_structures() {
// Test GetStatsRequest
+1
View File
@@ -268,6 +268,7 @@ pub async fn create_connector_by_url(
IpScheme::FakeTcp => tunnel::fake_tcp::FakeTcpTunnelConnector::new(url).boxed(),
};
connector.set_resolved_addr(resolved_addr.addr);
connector.set_socket_mark(global_ctx.config.get_flags().socket_mark);
if global_ctx.config.get_flags().bind_device {
set_bind_addr_for_peer_connector(
&mut connector,
+61 -13
View File
@@ -533,6 +533,17 @@ struct NetworkOptions {
)]
bind_device: Option<bool>,
// SO_MARK (fwmark) is a Linux-family kernel feature. Gate the flag out
// entirely on other targets so users on Windows/macOS/BSD don't see a
// `--socket-mark` they can't act on.
#[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))]
#[arg(
long,
env = "ET_SOCKET_MARK",
help = t!("core_clap.socket_mark").to_string()
)]
socket_mark: Option<u32>,
#[arg(
long,
env = "ET_ENABLE_KCP_PROXY",
@@ -880,17 +891,20 @@ impl NetworkOptions {
}
let old_ns = cfg.get_network_identity();
let network_name = self.network_name.clone().unwrap_or(old_ns.network_name);
let network_name = self
.network_name
.clone()
.unwrap_or_else(|| old_ns.network_name.clone());
if self.credential.is_some() {
// Credential mode: no network_secret, authenticate via credential keypair
cfg.set_network_identity(NetworkIdentity::new_credential(network_name));
} else {
let network_secret = self
.network_secret
.clone()
.unwrap_or(old_ns.network_secret.unwrap_or_default());
} else if let Some(network_secret) = &self.network_secret {
cfg.set_network_identity(NetworkIdentity::new(network_name, network_secret.clone()));
} else if let Some(network_secret) = old_ns.network_secret {
cfg.set_network_identity(NetworkIdentity::new(network_name, network_secret));
} else {
cfg.set_network_identity(NetworkIdentity::new_credential(network_name));
}
if let Some(dhcp) = self.dhcp {
@@ -1126,6 +1140,10 @@ impl NetworkOptions {
.into();
}
f.bind_device = self.bind_device.unwrap_or(f.bind_device);
#[cfg(any(target_os = "android", target_os = "fuchsia", target_os = "linux"))]
{
f.socket_mark = self.socket_mark.or(f.socket_mark);
}
f.enable_kcp_proxy = self.enable_kcp_proxy.unwrap_or(f.enable_kcp_proxy);
f.disable_kcp_input = self.disable_kcp_input.unwrap_or(f.disable_kcp_input);
f.enable_quic_proxy = self.enable_quic_proxy.unwrap_or(f.enable_quic_proxy);
@@ -1596,7 +1614,7 @@ pub async fn main() -> ExitCode {
// Verify configurations
if cli.check_config {
if let Err(error) = validate_config(&cli).await {
log::error!(?error, "Config validation failed");
log::error!(%error, "Config validation failed");
return ExitCode::FAILURE;
} else {
return ExitCode::SUCCESS;
@@ -1606,7 +1624,7 @@ pub async fn main() -> ExitCode {
let mut ret_code = 0;
if let Err(error) = run_main(cli).await {
log::error!(?error);
log::error!(%error);
ret_code = 1;
}
@@ -1626,12 +1644,13 @@ async fn validate_config(cli: &Cli) -> anyhow::Result<()> {
for config_file in config_files {
if config_file == &PathBuf::from("-") {
let mut stdin = String::new();
_ = tokio::io::stdin().read_to_string(&mut stdin).await?;
TomlConfigLoader::new_from_str(stdin.as_str())
.with_context(|| "config source: stdin")?;
_ = tokio::io::stdin()
.read_to_string(&mut stdin)
.await
.context("failed to read config from stdin")?;
TomlConfigLoader::new_from_str_with_source("stdin", stdin.as_str())?;
} else {
TomlConfigLoader::new(config_file)
.with_context(|| format!("config source: {:?}", config_file))?;
TomlConfigLoader::new(config_file)?;
};
}
@@ -1724,4 +1743,33 @@ mod tests {
);
}
}
#[test]
fn test_network_options_merge_preserves_credential_identity() {
let cfg = TomlConfigLoader::new_from_str(
r#"
[network_identity]
network_name = "credential-network"
network_secret = ""
[secure_mode]
enabled = true
"#,
)
.unwrap();
assert_eq!(cfg.get_network_identity().network_secret, None);
NetworkOptions {
hostname: Some("override-host".to_string()),
..Default::default()
}
.merge_into(&cfg)
.unwrap();
let identity = cfg.get_network_identity();
assert_eq!(identity.network_name, "credential-network");
assert_eq!(identity.network_secret, None);
assert_eq!(identity.network_secret_digest, None);
assert_eq!(cfg.get_hostname(), "override-host");
}
}
+3
View File
@@ -23,6 +23,9 @@ pub static malloc_conf: &[u8] = b"retain:false\0";
rust_i18n::i18n!("locales", fallback = "en");
// The easytier-core CLI intentionally uses a single-thread runtime. The
// `multi_thread` flag only affects launcher-based deployments (GUI / mobile /
// web / Windows service); see launcher.rs:223.
#[tokio::main(flavor = "current_thread")]
async fn main() -> std::process::ExitCode {
core::main().await
+53 -1
View File
@@ -133,7 +133,7 @@ impl AsyncUdpSocket for QuicSocket {
unsafe {
copy_nonoverlapping(
chunk.as_ptr(),
payload.as_mut_ptr().add(self.margins.header),
payload.chunk_mut().as_mut_ptr().add(self.margins.header),
len,
);
payload.advance_mut(len + self.margins.len());
@@ -1383,4 +1383,56 @@ mod tests {
Ok(())
}
#[tokio::test]
async fn test_gso() {
let margins = PacketMargins {
header: 20,
trailer: 25,
};
let (tx, rx) = channel(10);
let socket = QuicSocket {
addr: "127.0.0.1:0".parse().unwrap(),
rx: AtomicRefCell::new(rx),
tx,
margins,
};
let total_len = 3000;
let segment_size = 1000;
let mut contents = Vec::with_capacity(total_len);
contents.extend(vec![1u8; 1000]);
contents.extend(vec![2u8; 1000]);
contents.extend(vec![3u8; 1000]);
let transmit = Transmit {
destination: "127.0.0.1:8000".parse().unwrap(),
ecn: None,
contents: &contents,
segment_size: Some(segment_size),
src_ip: None,
};
socket.try_send(&transmit).unwrap();
let mut rx = socket.rx.into_inner();
let packet = rx.recv().await.unwrap();
let actual_segment_size = segment_size + margins.len();
let payload = packet.payload;
let chunk1_start = margins.header;
let chunk1_data = &payload[chunk1_start..chunk1_start + segment_size];
assert_eq!(chunk1_data[0], 1u8, "Chunk 1 corrupted");
let chunk2_start = actual_segment_size + margins.header;
let chunk2_data = &payload[chunk2_start..chunk2_start + segment_size];
assert_eq!(chunk2_data[0], 2u8, "Chunk 2 corrupted");
let chunk3_start = actual_segment_size * 2 + margins.header;
let chunk3_data = &payload[chunk3_start..chunk3_start + segment_size];
assert_eq!(chunk3_data[0], 3u8, "Chunk 3 corrupted");
}
}
+57 -3
View File
@@ -52,6 +52,12 @@ use crate::{
peers::{PeerPacketFilter, peer_manager::PeerManager},
};
#[cfg(feature = "ffi-dataplane")]
mod dataplane;
#[cfg(feature = "ffi-dataplane")]
pub use dataplane::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket};
enum SocksUdpSocket {
UdpSocket(Arc<tokio::net::UdpSocket>),
SmolUdpSocket(super::tokio_smoltcp::UdpSocket),
@@ -136,11 +142,17 @@ impl AsyncWrite for SocksTcpStream {
enum Socks5EntryData {
Tcp(TcpListener), // hold a binded socket to hold the tcp port
#[cfg(feature = "ffi-dataplane")]
// a data-plane routing entry that owns no resource. the entry_type in the
// key distinguishes a listen route from an actively outbound route.
DataPlaneRoute,
Udp((Arc<SocksUdpSocket>, UdpClientKey)), // hold the socket to send data to dst
}
const UDP_ENTRY: u8 = 1;
const TCP_ENTRY: u8 = 2;
#[cfg(feature = "ffi-dataplane")]
const TCP_LISTEN_ENTRY: u8 = 3;
#[derive(Debug, Eq, PartialEq, Hash, Clone)]
struct Socks5Entry {
@@ -477,6 +489,11 @@ pub struct Socks5Server {
kcp_endpoint: Mutex<Option<Weak<KcpEndpoint>>>,
socks5_enabled: Arc<AtomicBool>,
#[cfg(feature = "ffi-dataplane")]
data_plane_refs: Arc<AtomicUsize>,
// Tracks whether the smoltcp `net` is ready for data-plane callers.
#[cfg(feature = "ffi-dataplane")]
data_plane_net_ready: tokio::sync::watch::Sender<bool>,
cancel_tokens: Arc<DashMap<PortForwardConfig, DropGuard>>,
port_forward_list_change_notifier: Arc<Notify>,
entry_count: Arc<AtomicUsize>,
@@ -508,14 +525,27 @@ impl PeerPacketFilter for Socks5Server {
let Some(tcp_packet) = TcpPacket::new(ipv4.payload()) else {
return Some(packet);
};
Socks5Entry {
let entry = Socks5Entry {
dst: SocketAddr::new(ipv4.get_source().into(), tcp_packet.get_source()),
src: SocketAddr::new(
ipv4.get_destination().into(),
tcp_packet.get_destination(),
),
entry_type: TCP_ENTRY,
}
};
#[cfg(feature = "ffi-dataplane")]
let entry = if self.entries.contains_key(&entry) {
// Case 1: it is an established connection that has an exactly matched inbound.
entry
} else {
// Case 2: it could be a new TCP SYN packet that has not been accepted.
Socks5Entry {
src: entry.src,
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
entry_type: TCP_LISTEN_ENTRY,
}
};
entry
}
IpNextHeaderProtocols::Udp => {
@@ -591,6 +621,10 @@ impl Socks5Server {
kcp_endpoint: Mutex::new(None),
socks5_enabled: Arc::new(AtomicBool::new(false)),
#[cfg(feature = "ffi-dataplane")]
data_plane_refs: Arc::new(AtomicUsize::new(0)),
#[cfg(feature = "ffi-dataplane")]
data_plane_net_ready: tokio::sync::watch::channel(false).0,
cancel_tokens: Arc::new(DashMap::new()),
port_forward_list_change_notifier: Arc::new(Notify::new()),
entry_count: Arc::new(AtomicUsize::new(0)),
@@ -608,11 +642,25 @@ impl Socks5Server {
let cancel_tokens = self.cancel_tokens.clone();
let port_forward_list_change_notifier = self.port_forward_list_change_notifier.clone();
let socks5_enabled = self.socks5_enabled.clone();
#[cfg(feature = "ffi-dataplane")]
let data_plane_refs = self.data_plane_refs.clone();
#[cfg(feature = "ffi-dataplane")]
let data_plane_net_ready = self.data_plane_net_ready.clone();
self.tasks.lock().unwrap().spawn(async move {
let mut prev_ipv4 = None;
loop {
if cancel_tokens.is_empty() && !socks5_enabled.load(Ordering::Relaxed) {
#[cfg(feature = "ffi-dataplane")]
let data_plane_active = data_plane_refs.load(Ordering::Relaxed) > 0;
#[cfg(not(feature = "ffi-dataplane"))]
let data_plane_active = false;
if cancel_tokens.is_empty()
&& !socks5_enabled.load(Ordering::Relaxed)
&& !data_plane_active
{
let _ = net.lock().await.take();
#[cfg(feature = "ffi-dataplane")]
let _ = data_plane_net_ready.send_replace(false);
port_forward_list_change_notifier.notified().await;
continue;
}
@@ -637,8 +685,14 @@ impl Socks5Server {
packet_recv.clone(),
entries.clone(),
));
// Wake any data-plane callers waiting in
// `wait_data_plane_net` for the smoltcp net to appear.
#[cfg(feature = "ffi-dataplane")]
let _ = data_plane_net_ready.send_replace(true);
} else {
let _ = net.lock().await.take();
#[cfg(feature = "ffi-dataplane")]
let _ = data_plane_net_ready.send_replace(false);
}
}
+535
View File
@@ -0,0 +1,535 @@
//! Data-plane access built on top of the `Socks5Server` smoltcp stack.
//!
//! This module exposes TCP streams and UDP sockets (mainly for FFI callers that
//! send traffic through EasyTier without creating OS-level proxy listeners).
//!
//! Typical usage:
//!
//! ```ignore
//! let instance = Instance::new(cfg);
//! instance.run().await?;
//! let socks5_server = instance.get_socks5_server();
//!
//! let socket = socks5_server.data_plane_udp_bind(local_port, timeout).await?;
//! socket.send_to(buf, peer_addr).await?;
//! ```
use std::{
net::{IpAddr, Ipv4Addr, SocketAddr},
pin::Pin,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll},
time::{Duration, Instant},
};
use anyhow::Context as _;
use dashmap::mapref::entry::Entry;
use tokio::io::{AsyncRead, AsyncWrite};
use crate::{common::error::Error, gateway::fast_socks5::server::AsyncTcpConnector};
use super::{
Socks5AutoConnector, Socks5Entry, Socks5EntryData, Socks5EntrySet, Socks5Server,
SocksTcpStream, SocksUdpSocket, TCP_ENTRY, TCP_LISTEN_ENTRY, UDP_ENTRY, UdpClientKey,
};
use crate::gateway::tokio_smoltcp::{Net, TcpListener};
struct DataPlaneRef {
refs: Arc<AtomicUsize>,
notifier: Arc<tokio::sync::Notify>,
}
/// A route-table entry whose lifetime is tied to this value: constructing it
/// reserves the route and bumps the active-entry count, dropping it removes the
/// route and drops the count back.
struct OwnedRouteEntry {
entries: Socks5EntrySet,
entry_count: Arc<AtomicUsize>,
entry: Socks5Entry,
}
impl OwnedRouteEntry {
/// Inserts the route, replacing any existing entry for the same key.
fn register(
entries: Socks5EntrySet,
entry_count: Arc<AtomicUsize>,
entry: Socks5Entry,
) -> Self {
if entries
.insert(entry.clone(), Socks5EntryData::DataPlaneRoute)
.is_none()
{
entry_count.fetch_add(1, Ordering::Relaxed);
}
Self {
entries,
entry_count,
entry,
}
}
/// Inserts the route only if the key is free, returning `None` on conflict.
fn try_register(
entries: Socks5EntrySet,
entry_count: Arc<AtomicUsize>,
entry: Socks5Entry,
) -> Option<Self> {
match entries.entry(entry.clone()) {
Entry::Occupied(_) => return None,
Entry::Vacant(vacant) => {
vacant.insert(Socks5EntryData::DataPlaneRoute);
entry_count.fetch_add(1, Ordering::Relaxed);
}
}
Some(Self {
entries,
entry_count,
entry,
})
}
}
impl Drop for OwnedRouteEntry {
fn drop(&mut self) {
if self.entries.remove(&self.entry).is_some() {
self.entry_count.fetch_sub(1, Ordering::Relaxed);
}
}
}
/// Tracks how an established data-plane TCP stream keeps its inbound route alive.
///
/// The two variants capture the intrinsic asymmetry between the connect and
/// accept paths. An outbound stream reserved a source port through the
/// [`Socks5AutoConnector`], which owns the matching route entry and clears it on
/// drop. An accepted stream instead inherits its port and peer from the
/// listener, so it carries merely an [`OwnedRouteEntry`].
enum DataPlaneTcpStreamRoute {
Outbound(Socks5AutoConnector),
Accepted(OwnedRouteEntry),
}
/// A TCP stream created by the data plane API.
/// Can be either an actively requested outbound connection or an outbound request accepted from a TCP listener.
pub struct DataPlaneTcpStream {
stream: SocksTcpStream,
local_addr: SocketAddr,
_data_plane_ref: DataPlaneRef,
_route: DataPlaneTcpStreamRoute,
}
/// A TCP listener created by the data plane API.
/// It accepts inbound connections and produces [`DataPlaneTcpStream`]s.
pub struct DataPlaneTcpListener {
listener: TcpListener,
local_addr: SocketAddr,
entries: Socks5EntrySet,
entry_count: Arc<AtomicUsize>,
_listen_route: OwnedRouteEntry,
_data_plane_ref: DataPlaneRef,
}
pub struct DataPlaneUdpSocket {
socket: Arc<SocksUdpSocket>,
entries: Socks5EntrySet,
entry_count: Arc<AtomicUsize>,
local_addr: SocketAddr,
_data_plane_ref: DataPlaneRef,
}
impl Drop for DataPlaneRef {
fn drop(&mut self) {
if self.refs.fetch_sub(1, Ordering::Relaxed) == 1 {
self.notifier.notify_one();
}
}
}
impl DataPlaneTcpStream {
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
}
impl DataPlaneTcpListener {
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub async fn accept(&mut self) -> Result<(DataPlaneTcpStream, SocketAddr), std::io::Error> {
let (stream, peer_addr) = self.listener.accept().await?;
let local_addr = stream.local_addr()?;
let route = OwnedRouteEntry::register(
self.entries.clone(),
self.entry_count.clone(),
Socks5Entry {
src: local_addr,
dst: peer_addr,
entry_type: TCP_ENTRY,
},
);
let accepted = DataPlaneTcpStream {
stream: SocksTcpStream::SmolTcp(stream),
local_addr,
_data_plane_ref: self._data_plane_ref.clone(),
_route: DataPlaneTcpStreamRoute::Accepted(route),
};
Ok((accepted, peer_addr))
}
}
impl AsyncRead for DataPlaneTcpStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().stream).poll_read(cx, buf)
}
}
impl AsyncWrite for DataPlaneTcpStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
Pin::new(&mut self.get_mut().stream).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.get_mut().stream).poll_flush(cx)
}
fn poll_shutdown(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
Pin::new(&mut self.get_mut().stream).poll_shutdown(cx)
}
}
impl DataPlaneUdpSocket {
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub async fn send_to(&self, buf: &[u8], addr: SocketAddr) -> Result<usize, std::io::Error> {
let key = Socks5Entry {
src: self.local_addr,
dst: addr,
entry_type: UDP_ENTRY,
};
if let Entry::Vacant(entry) = self.entries.entry(key) {
entry.insert(Socks5EntryData::Udp((
self.socket.clone(),
UdpClientKey {
client_addr: self.local_addr,
dst_addr: addr,
},
)));
self.entry_count.fetch_add(1, Ordering::Relaxed);
}
self.socket.send_to(buf, addr).await
}
pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, SocketAddr), std::io::Error> {
self.socket.recv_from(buf).await
}
}
impl Drop for DataPlaneUdpSocket {
fn drop(&mut self) {
self.entries.retain(|_, data| match data {
Socks5EntryData::Udp((socket, _)) if Arc::ptr_eq(socket, &self.socket) => {
self.entry_count.fetch_sub(1, Ordering::Relaxed);
false
}
_ => true,
});
}
}
impl Clone for DataPlaneRef {
fn clone(&self) -> Self {
self.refs.fetch_add(1, Ordering::Relaxed);
Self {
refs: self.refs.clone(),
notifier: self.notifier.clone(),
}
}
}
impl Socks5Server {
fn acquire_data_plane_ref(&self) -> DataPlaneRef {
self.data_plane_refs.fetch_add(1, Ordering::Relaxed);
self.port_forward_list_change_notifier.notify_one();
DataPlaneRef {
refs: self.data_plane_refs.clone(),
notifier: self.port_forward_list_change_notifier.clone(),
}
}
async fn wait_data_plane_net(
&self,
deadline: Instant,
) -> Result<(cidr::Ipv4Inet, Arc<Net>), Error> {
let mut ready = self.data_plane_net_ready.subscribe();
loop {
if let Some(net) = self
.net
.lock()
.await
.as_ref()
.map(|net| (net.ipv4_addr, net.smoltcp_net.clone()))
{
return Ok(net);
}
let now = Instant::now();
if now >= deadline {
return Err(anyhow::anyhow!("data plane net is not ready").into());
}
let _ = tokio::time::timeout(deadline - now, ready.wait_for(|ready| *ready)).await;
}
}
pub async fn data_plane_tcp_connect(
&self,
dst_addr: SocketAddr,
timeout: Duration,
) -> Result<DataPlaneTcpStream, Error> {
let data_plane_ref = self.acquire_data_plane_ref();
let deadline = Instant::now() + timeout;
let (ipv4_addr, smoltcp_net) = self.wait_data_plane_net(deadline).await?;
// FIXME: This is the data-plane source address reserved for route
// matching. `Socks5AutoConnector` may fall back to direct TCP for
// non-virtual destinations, so this is not always the OS socket's
// local address.
let local_port = smoltcp_net.get_port();
let local_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port);
let connector = Socks5AutoConnector {
#[cfg(feature = "kcp")]
kcp_endpoint: self.kcp_endpoint.lock().await.clone(),
peer_mgr: self.peer_manager.clone(),
entries: self.entries.clone(),
smoltcp_net: Some(smoltcp_net),
src_addr: local_addr,
entry_count: self.entry_count.clone(),
inner_connector: parking_lot::Mutex::new(None),
};
let remaining = deadline.saturating_duration_since(Instant::now());
let inner_timeout_s = remaining.as_secs().saturating_add(1);
let stream =
tokio::time::timeout(remaining, connector.tcp_connect(dst_addr, inner_timeout_s))
.await
.with_context(|| "data plane tcp connect timeout")?
.map_err(anyhow::Error::from)?;
Ok(DataPlaneTcpStream {
stream,
local_addr,
_data_plane_ref: data_plane_ref,
_route: DataPlaneTcpStreamRoute::Outbound(connector),
})
}
pub async fn data_plane_tcp_bind(
&self,
local_port: u16,
timeout: Duration,
) -> Result<DataPlaneTcpListener, Error> {
let data_plane_ref = self.acquire_data_plane_ref();
let deadline = Instant::now() + timeout;
let (ipv4_addr, smoltcp_net) = self.wait_data_plane_net(deadline).await?;
let bind_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port);
let listener = smoltcp_net.tcp_bind(bind_addr).await?;
let local_addr = listener.local_addr()?;
let listen_route = OwnedRouteEntry::try_register(
self.entries.clone(),
self.entry_count.clone(),
Socks5Entry {
src: local_addr,
dst: SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
entry_type: TCP_LISTEN_ENTRY,
},
)
.ok_or_else(|| anyhow::anyhow!("data plane tcp listener already exists"))?;
Ok(DataPlaneTcpListener {
listener,
local_addr,
entries: self.entries.clone(),
entry_count: self.entry_count.clone(),
_listen_route: listen_route,
_data_plane_ref: data_plane_ref,
})
}
pub async fn data_plane_udp_bind(
&self,
local_port: u16,
timeout: Duration,
) -> Result<DataPlaneUdpSocket, Error> {
let data_plane_ref = self.acquire_data_plane_ref();
let deadline = Instant::now() + timeout;
let (ipv4_addr, smoltcp_net) = self.wait_data_plane_net(deadline).await?;
let bind_addr = SocketAddr::new(IpAddr::V4(ipv4_addr.address()), local_port);
let smol = smoltcp_net.udp_bind(bind_addr).await?;
let local_addr = smol.local_addr()?;
let socket = Arc::new(SocksUdpSocket::SmolUdpSocket(smol));
Ok(DataPlaneUdpSocket {
socket,
entries: self.entries.clone(),
entry_count: self.entry_count.clone(),
local_addr,
_data_plane_ref: data_plane_ref,
})
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use super::Socks5Server;
use crate::peers::peer_manager::PeerManager;
use crate::peers::tests::{connect_peer_manager, create_mock_peer_manager};
use crate::tunnel::common::tests::wait_for_condition;
/// A peer and its data-plane server. `Socks5Server` only holds a `Weak`
/// reference to the `PeerManager`, so the manager must be kept alive by the
/// test for the server's smoltcp <-> peer routing to work.
struct Endpoint {
_peer: std::sync::Arc<PeerManager>,
server: std::sync::Arc<Socks5Server>,
ip: cidr::Ipv4Inet,
}
/// Brings up two peers connected by a ring tunnel, each with a virtual IPv4
/// and a running `Socks5Server`, and waits until the route to `b`'s IPv4 is
/// visible from `a`. `run(None)` leaves the kcp endpoint unset, so the
/// connect path goes through smoltcp, matching the listener side under test.
async fn setup_pair() -> (Endpoint, Endpoint) {
let a = create_mock_peer_manager().await;
let b = create_mock_peer_manager().await;
connect_peer_manager(a.clone(), b.clone()).await;
let a_ip: cidr::Ipv4Inet = "10.126.126.1/24".parse().unwrap();
let b_ip: cidr::Ipv4Inet = "10.126.126.2/24".parse().unwrap();
a.get_global_ctx().set_ipv4(Some(a_ip));
b.get_global_ctx().set_ipv4(Some(b_ip));
let server_a = Socks5Server::new(a.get_global_ctx(), a.clone(), None);
let server_b = Socks5Server::new(b.get_global_ctx(), b.clone(), None);
server_a.run(None).await.unwrap();
server_b.run(None).await.unwrap();
wait_for_condition(
|| async {
a.get_route()
.get_peer_id_by_ipv4(&b_ip.address())
.await
.is_some()
},
Duration::from_secs(10),
)
.await;
(
Endpoint {
_peer: a,
server: server_a,
ip: a_ip,
},
Endpoint {
_peer: b,
server: server_b,
ip: b_ip,
},
)
}
#[tokio::test]
async fn data_plane_tcp_pingpong() {
let (ep_a, ep_b) = setup_pair().await;
let (server_a, server_b, b_ip) = (ep_a.server, ep_b.server, ep_b.ip);
let timeout = Duration::from_secs(10);
let mut listener = server_b.data_plane_tcp_bind(0, timeout).await.unwrap();
let listen_addr =
std::net::SocketAddr::new(b_ip.address().into(), listener.local_addr().port());
let accept = tokio::spawn(async move {
let (mut stream, _peer) = listener.accept().await.unwrap();
let mut buf = [0u8; 4];
stream.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"ping");
stream.write_all(b"pong").await.unwrap();
stream.flush().await.unwrap();
// Hold the listener and stream until the client has read the reply.
tokio::time::sleep(Duration::from_secs(1)).await;
});
let mut client = server_a
.data_plane_tcp_connect(listen_addr, timeout)
.await
.unwrap();
client.write_all(b"ping").await.unwrap();
client.flush().await.unwrap();
let mut buf = [0u8; 4];
client.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"pong");
accept.await.unwrap();
}
#[tokio::test]
async fn data_plane_udp_pingpong() {
let (ep_a, ep_b) = setup_pair().await;
let (server_a, a_ip, server_b, b_ip) = (ep_a.server, ep_a.ip, ep_b.server, ep_b.ip);
let timeout = Duration::from_secs(10);
let sock_a = server_a.data_plane_udp_bind(0, timeout).await.unwrap();
let sock_b = server_b.data_plane_udp_bind(0, timeout).await.unwrap();
let addr_a = std::net::SocketAddr::new(a_ip.address().into(), sock_a.local_addr().port());
let addr_b = std::net::SocketAddr::new(b_ip.address().into(), sock_b.local_addr().port());
// UDP data-plane routes are connected-style: a socket only accepts
// inbound datagrams from a peer it has already sent to, because the
// route entry is registered by `send_to`. Prime b's route toward a so
// the upcoming ping is routed instead of dropped at b's packet filter.
// This datagram is dropped at a (a has no route yet) and is not awaited.
sock_b.send_to(b"warmup", addr_a).await.unwrap();
sock_a.send_to(b"ping", addr_b).await.unwrap();
let mut buf = [0u8; 16];
let (n, from) = tokio::time::timeout(timeout, sock_b.recv_from(&mut buf))
.await
.expect("recv ping timed out")
.unwrap();
assert_eq!(&buf[..n], b"ping");
assert_eq!(from, addr_a);
sock_b.send_to(b"pong", addr_a).await.unwrap();
// a may also receive the stray warmup datagram (it arrives once a has
// registered its route by sending the ping above), so skip anything
// that is not the reply.
loop {
let (n, from) = tokio::time::timeout(timeout, sock_a.recv_from(&mut buf))
.await
.expect("recv pong timed out")
.unwrap();
if &buf[..n] == b"pong" {
assert_eq!(from, addr_b);
break;
}
}
}
}
+334 -7
View File
@@ -40,6 +40,16 @@ use crate::{
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
struct UdpNatKey {
src_socket: SocketAddr,
dst_socket: SocketAddr,
}
impl UdpNatKey {
fn new(src_socket: SocketAddr, dst_socket: SocketAddr) -> Self {
Self {
src_socket,
dst_socket,
}
}
}
#[derive(Debug)]
@@ -204,13 +214,23 @@ impl UdpNatEntry {
self_clone.mark_active();
if src_v4.ip().is_loopback() {
src_v4.set_ip(virtual_ipv4);
let has_mapped_dst = real_ipv4 != mapped_ipv4;
let mut reply_src_ip = *src_v4.ip();
// Preserve the existing priority for proxy rules that expose a
// real loopback address as a mapped address. Other loopback
// replies come from local delivery to 127.0.0.1 for the local
// virtual IP and may need the mapped rewrite below.
if has_mapped_dst && reply_src_ip == real_ipv4 {
reply_src_ip = mapped_ipv4;
} else if reply_src_ip.is_loopback() {
reply_src_ip = virtual_ipv4;
}
if *src_v4.ip() == real_ipv4 {
src_v4.set_ip(mapped_ipv4);
if has_mapped_dst && reply_src_ip == real_ipv4 {
reply_src_ip = mapped_ipv4;
}
src_v4.set_ip(reply_src_ip);
let Ok(_) = Self::compose_ipv4_packet(
&self_clone,
@@ -321,9 +341,10 @@ impl UdpProxy {
"udp nat packet request received"
);
let nat_key = UdpNatKey {
src_socket: SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()),
};
let nat_key = UdpNatKey::new(
SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()),
SocketAddr::new(ipv4.get_destination().into(), udp_packet.get_destination()),
);
let nat_entry = self
.nat_table
.entry(nat_key)
@@ -487,3 +508,309 @@ impl Drop for UdpProxy {
}
}
}
#[cfg(test)]
mod tests {
use std::{
net::{Ipv4Addr, SocketAddr},
sync::Arc,
time::Duration,
};
use pnet::packet::{
MutablePacket, Packet,
ip::IpNextHeaderProtocols,
ipv4::{self, Ipv4Packet, MutableIpv4Packet},
udp::{self, MutableUdpPacket, UdpPacket},
};
use tokio::{net::UdpSocket, sync::mpsc::Receiver, time::timeout};
use crate::{
common::{config::ConfigLoader, global_ctx::tests::get_mock_global_ctx},
peers::{
create_packet_recv_chan,
peer_manager::{PeerManager, RouteAlgoType},
},
tunnel::packet_def::{PacketType, ZCPacket},
};
use super::UdpProxy;
fn build_udp_proxy_packet(
src_ip: Ipv4Addr,
src_port: u16,
dst_socket: SocketAddr,
payload: &[u8],
) -> ZCPacket {
let SocketAddr::V4(dst_socket) = dst_socket else {
panic!("test only builds IPv4 UDP packets");
};
let dst_ip = *dst_socket.ip();
let mut packet = vec![0; 20 + 8 + payload.len()];
let packet_len = packet.len() as u16;
{
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
ipv4_packet.set_version(4);
ipv4_packet.set_header_length(5);
ipv4_packet.set_total_length(packet_len);
ipv4_packet.set_ttl(64);
ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp);
ipv4_packet.set_source(src_ip);
ipv4_packet.set_destination(dst_ip);
}
{
let mut udp_packet = MutableUdpPacket::new(&mut packet[20..]).unwrap();
udp_packet.set_source(src_port);
udp_packet.set_destination(dst_socket.port());
udp_packet.set_length((8 + payload.len()) as u16);
udp_packet.payload_mut().copy_from_slice(payload);
udp_packet.set_checksum(udp::ipv4_checksum(
&udp_packet.to_immutable(),
&src_ip,
&dst_ip,
));
}
{
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
ipv4_packet.set_checksum(ipv4::checksum(&ipv4_packet.to_immutable()));
}
let mut packet = ZCPacket::new_with_payload(&packet);
packet.fill_peer_manager_hdr(1009867077, 3831440917, PacketType::Data as u8);
packet
}
async fn wait_proxy_cidr_loaded(proxy: &UdpProxy) {
timeout(Duration::from_secs(1), async {
while proxy.cidr_set.is_empty() {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
}
async fn recv_payload(socket: &UdpSocket) -> (Vec<u8>, SocketAddr) {
let mut buf = [0; 64];
let (len, addr) = timeout(Duration::from_secs(1), socket.recv_from(&mut buf))
.await
.unwrap()
.unwrap();
(buf[..len].to_vec(), addr)
}
async fn recv_response_packet(receiver: &mut Receiver<ZCPacket>) -> ZCPacket {
timeout(Duration::from_secs(1), receiver.recv())
.await
.unwrap()
.unwrap()
}
fn assert_udp_response(
packet: ZCPacket,
src_socket: SocketAddr,
dst_ip: Ipv4Addr,
dst_port: u16,
payload: &[u8],
) {
let SocketAddr::V4(src_socket) = src_socket else {
panic!("test only checks IPv4 UDP packets");
};
let ipv4_packet = Ipv4Packet::new(packet.payload()).unwrap();
assert_eq!(ipv4_packet.get_source(), *src_socket.ip());
assert_eq!(ipv4_packet.get_destination(), dst_ip);
let udp_packet = UdpPacket::new(ipv4_packet.payload()).unwrap();
assert_eq!(udp_packet.get_source(), src_socket.port());
assert_eq!(udp_packet.get_destination(), dst_port);
assert_eq!(udp_packet.payload(), payload);
}
async fn stop_nat_entries(proxy: &UdpProxy) {
let nat_socket_addrs = proxy
.nat_table
.iter()
.filter_map(|entry| {
entry
.socket
.as_ref()
.and_then(|socket| socket.local_addr().ok())
.map(|addr| SocketAddr::from((Ipv4Addr::LOCALHOST, addr.port())))
})
.collect::<Vec<_>>();
for entry in proxy.nat_table.iter() {
entry.stop();
}
let wake_socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
for addr in nat_socket_addrs {
let _ = wake_socket.send_to(b"wake", addr).await;
}
}
#[tokio::test]
async fn udp_proxy_rewrites_unmapped_loopback_reply_to_virtual_ip() {
let global_ctx = get_mock_global_ctx();
global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap()));
global_ctx
.config
.add_proxy_cidr("127.0.0.1/32".parse().unwrap(), None)
.unwrap();
let (packet_sender, _packet_receiver) = create_packet_recv_chan();
let peer_manager = Arc::new(PeerManager::new(
RouteAlgoType::Ospf,
global_ctx.clone(),
packet_sender,
));
let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap();
wait_proxy_cidr_loaded(&proxy).await;
let mut response_receiver = proxy.receiver.lock().await.take().unwrap();
let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let real_dst_port = real_dst.local_addr().unwrap().port();
let dst_socket = SocketAddr::from((Ipv4Addr::LOCALHOST, real_dst_port));
let src_ip = Ipv4Addr::new(10, 144, 144, 206);
let src_port = 53864;
let packet = build_udp_proxy_packet(src_ip, src_port, dst_socket, b"request");
assert!(proxy.try_handle_packet(&packet).await.is_some());
let (payload, nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"request");
real_dst.send_to(b"reply", nat_socket).await.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
SocketAddr::from((Ipv4Addr::new(10, 144, 144, 204), real_dst_port)),
src_ip,
src_port,
b"reply",
);
stop_nat_entries(&proxy).await;
}
#[tokio::test]
async fn udp_proxy_maps_local_virtual_destination_reply_to_mapped_source() {
let global_ctx = get_mock_global_ctx();
global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap()));
global_ctx
.config
.add_proxy_cidr(
"10.144.144.204/32".parse().unwrap(),
Some("10.10.10.3/32".parse().unwrap()),
)
.unwrap();
let (packet_sender, _packet_receiver) = create_packet_recv_chan();
let peer_manager = Arc::new(PeerManager::new(
RouteAlgoType::Ospf,
global_ctx.clone(),
packet_sender,
));
let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap();
wait_proxy_cidr_loaded(&proxy).await;
let mut response_receiver = proxy.receiver.lock().await.take().unwrap();
let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let real_dst_port = real_dst.local_addr().unwrap().port();
let mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 3), real_dst_port));
let src_ip = Ipv4Addr::new(10, 144, 144, 206);
let src_port = 53864;
let packet = build_udp_proxy_packet(src_ip, src_port, mapped_dst, b"request");
assert!(proxy.try_handle_packet(&packet).await.is_some());
let (payload, nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"request");
real_dst.send_to(b"reply", nat_socket).await.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
mapped_dst,
src_ip,
src_port,
b"reply",
);
stop_nat_entries(&proxy).await;
}
#[tokio::test]
async fn udp_proxy_separates_same_source_port_to_multiple_mapped_destinations() {
let global_ctx = get_mock_global_ctx();
global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap()));
global_ctx
.config
.add_proxy_cidr(
"127.0.0.1/32".parse().unwrap(),
Some("10.10.10.1/32".parse().unwrap()),
)
.unwrap();
global_ctx
.config
.add_proxy_cidr(
"127.0.0.1/32".parse().unwrap(),
Some("10.10.10.2/32".parse().unwrap()),
)
.unwrap();
let (packet_sender, _packet_receiver) = create_packet_recv_chan();
let peer_manager = Arc::new(PeerManager::new(
RouteAlgoType::Ospf,
global_ctx.clone(),
packet_sender,
));
let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap();
wait_proxy_cidr_loaded(&proxy).await;
let mut response_receiver = proxy.receiver.lock().await.take().unwrap();
let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let real_dst_port = real_dst.local_addr().unwrap().port();
let first_mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 1), real_dst_port));
let second_mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 2), real_dst_port));
let src_ip = Ipv4Addr::new(10, 144, 144, 206);
let src_port = 53864;
let first_packet = build_udp_proxy_packet(src_ip, src_port, first_mapped_dst, b"first");
assert!(proxy.try_handle_packet(&first_packet).await.is_some());
let (payload, first_nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"first");
let second_packet = build_udp_proxy_packet(src_ip, src_port, second_mapped_dst, b"second");
assert!(proxy.try_handle_packet(&second_packet).await.is_some());
let (payload, second_nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"second");
assert_eq!(proxy.nat_table.len(), 2);
real_dst
.send_to(b"first-reply", first_nat_socket)
.await
.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
first_mapped_dst,
src_ip,
src_port,
b"first-reply",
);
real_dst
.send_to(b"second-reply", second_nat_socket)
.await
.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
second_mapped_dst,
src_ip,
src_port,
b"second-reply",
);
stop_nat_entries(&proxy).await;
}
}
+5
View File
@@ -1143,6 +1143,11 @@ impl Instance {
self.peer_manager.clone()
}
#[cfg(feature = "ffi-dataplane")]
pub fn get_socks5_server(&self) -> Arc<Socks5Server> {
self.socks5_server.clone()
}
pub async fn close_peer_conn(
&mut self,
peer_id: PeerId,
+19 -4
View File
@@ -27,10 +27,20 @@ pub fn create_listener_by_url(
l: &url::Url,
global_ctx: ArcGlobalCtx,
) -> Result<Box<dyn TunnelListener>, Error> {
use crate::common::config::ConfigLoader;
let socket_mark = global_ctx.config.get_flags().socket_mark;
Ok(match l.try_into()? {
TunnelScheme::Ip(scheme) => match scheme {
IpScheme::Tcp => TcpTunnelListener::new(l.clone()).boxed(),
IpScheme::Udp => UdpTunnelListener::new(l.clone()).boxed(),
IpScheme::Tcp => {
let mut l = TcpTunnelListener::new(l.clone());
l.set_socket_mark(socket_mark);
l.boxed()
}
IpScheme::Udp => {
let mut l = UdpTunnelListener::new(l.clone());
l.set_socket_mark(socket_mark);
l.boxed()
}
#[cfg(feature = "wireguard")]
IpScheme::Wg => {
use crate::tunnel::wireguard::{WgConfig, WgTunnelListener};
@@ -39,15 +49,20 @@ pub fn create_listener_by_url(
&nid.network_name,
&nid.network_secret.unwrap_or_default(),
);
WgTunnelListener::new(l.clone(), wg_config).boxed()
let mut l = WgTunnelListener::new(l.clone(), wg_config);
l.set_socket_mark(socket_mark);
l.boxed()
}
#[cfg(feature = "quic")]
IpScheme::Quic => {
// QUIC reads socket_mark from global_ctx in QuicEndpointManager
tunnel::quic::QuicTunnelListener::new(l.clone(), global_ctx.clone()).boxed()
}
#[cfg(feature = "websocket")]
IpScheme::Ws | IpScheme::Wss => {
tunnel::websocket::WsTunnelListener::new(l.clone()).boxed()
let mut l = tunnel::websocket::WsTunnelListener::new(l.clone());
l.set_socket_mark(socket_mark);
l.boxed()
}
#[cfg(feature = "faketcp")]
IpScheme::FakeTcp => tunnel::fake_tcp::FakeTcpTunnelListener::new(l.clone()).boxed(),
+67
View File
@@ -1,3 +1,5 @@
#[cfg(feature = "ffi-dataplane")]
use crate::launcher::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket};
use dashmap::DashMap;
use std::fmt::{Display, Formatter};
use std::{collections::BTreeMap, path::PathBuf, sync::Arc};
@@ -32,6 +34,7 @@ pub struct NetworkInstanceManager {
instance_error_messages: Arc<DashMap<uuid::Uuid, String>>,
config_dir: Option<PathBuf>,
guard_counter: Arc<()>,
remote_mutation_lock: Arc<tokio::sync::Mutex<()>>,
}
impl Default for NetworkInstanceManager {
@@ -49,6 +52,7 @@ impl NetworkInstanceManager {
instance_error_messages: Arc::new(DashMap::new()),
config_dir: None,
guard_counter: Arc::new(()),
remote_mutation_lock: Arc::new(tokio::sync::Mutex::new(())),
}
}
@@ -57,6 +61,10 @@ impl NetworkInstanceManager {
self
}
pub fn remote_mutation_lock(&self) -> Arc<tokio::sync::Mutex<()>> {
self.remote_mutation_lock.clone()
}
fn start_instance_task(&self, instance_id: uuid::Uuid) -> Result<(), anyhow::Error> {
if tokio::runtime::Handle::try_current().is_err() {
return Err(anyhow::anyhow!(
@@ -171,6 +179,59 @@ impl NetworkInstanceManager {
tokio::runtime::Runtime::new()?.block_on(self.collect_network_infos())
}
#[cfg(feature = "ffi-dataplane")]
pub async fn data_plane_tcp_connect(
&self,
instance_id: &uuid::Uuid,
dst_addr: std::net::SocketAddr,
timeout: std::time::Duration,
) -> Result<DataPlaneTcpStream, anyhow::Error> {
let instance = self
.instance_map
.get(instance_id)
.ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?;
instance.data_plane_tcp_connect(dst_addr, timeout).await
}
#[cfg(feature = "ffi-dataplane")]
pub async fn data_plane_tcp_bind(
&self,
instance_id: &uuid::Uuid,
local_port: u16,
timeout: std::time::Duration,
) -> Result<DataPlaneTcpListener, anyhow::Error> {
let instance = self
.instance_map
.get(instance_id)
.ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?;
instance.data_plane_tcp_bind(local_port, timeout).await
}
#[cfg(feature = "ffi-dataplane")]
pub async fn data_plane_udp_bind(
&self,
instance_id: &uuid::Uuid,
local_port: u16,
timeout: std::time::Duration,
) -> Result<DataPlaneUdpSocket, anyhow::Error> {
let instance = self
.instance_map
.get(instance_id)
.ok_or_else(|| anyhow::anyhow!("instance {} not found", instance_id))?;
instance.data_plane_udp_bind(local_port, timeout).await
}
#[cfg(feature = "ffi-dataplane")]
pub fn data_plane_wait_runtime_handle(
&self,
instance_id: &uuid::Uuid,
timeout: std::time::Duration,
) -> Option<tokio::runtime::Handle> {
self.instance_map
.get(instance_id)
.and_then(|inst| inst.wait_runtime_handle(timeout))
}
pub async fn get_network_info(
&self,
instance_id: &uuid::Uuid,
@@ -217,6 +278,12 @@ impl NetworkInstanceManager {
.map(|instance| instance.value().get_config_file_control().clone())
}
pub fn get_instance_config(&self, instance_id: &uuid::Uuid) -> Option<TomlConfigLoader> {
self.instance_map
.get(instance_id)
.map(|instance| instance.value().get_config())
}
pub fn get_instance_network_config_source(
&self,
instance_id: &uuid::Uuid,
+141
View File
@@ -2,6 +2,10 @@ use crate::common::config::{
ConfigFileControl, ConfigSource, PortForwardConfig, parse_mapped_listener_urls,
process_secure_mode_cfg,
};
#[cfg(feature = "ffi-dataplane")]
use crate::gateway::socks5::Socks5Server;
#[cfg(feature = "ffi-dataplane")]
pub use crate::gateway::socks5::{DataPlaneTcpListener, DataPlaneTcpStream, DataPlaneUdpSocket};
use crate::proto::api::{self, manage};
use crate::proto::rpc_types::controller::BaseController;
use crate::rpc_service::InstanceRpcService;
@@ -45,6 +49,13 @@ struct EasyTierData {
tun_fd: (mpsc::Sender<TunFd>, Mutex<Option<mpsc::Receiver<TunFd>>>),
event_subscriber: RwLock<broadcast::Sender<GlobalCtxEvent>>,
instance_stop_notifier: Arc<tokio::sync::Notify>,
#[cfg(feature = "ffi-dataplane")]
data_plane: tokio::sync::watch::Sender<Option<Arc<Socks5Server>>>,
#[cfg(feature = "ffi-dataplane")]
runtime_handle: (
parking_lot::Mutex<Option<tokio::runtime::Handle>>,
parking_lot::Condvar,
),
}
impl Default for EasyTierData {
@@ -56,6 +67,10 @@ impl Default for EasyTierData {
events: RwLock::new(VecDeque::new()),
tun_fd: (sender, Mutex::new(Some(receiver))),
instance_stop_notifier: Arc::new(tokio::sync::Notify::new()),
#[cfg(feature = "ffi-dataplane")]
data_plane: tokio::sync::watch::channel(None).0,
#[cfg(feature = "ffi-dataplane")]
runtime_handle: (parking_lot::Mutex::new(None), parking_lot::Condvar::new()),
}
}
}
@@ -160,6 +175,10 @@ impl EasyTierLauncher {
instance.run().await?;
#[cfg(feature = "ffi-dataplane")]
data.data_plane
.send_replace(Some(instance.get_socks5_server()));
api_service
.write()
.unwrap()
@@ -214,6 +233,13 @@ impl EasyTierLauncher {
}
.unwrap();
#[cfg(feature = "ffi-dataplane")]
{
let (lock, cvar) = &data.runtime_handle;
*lock.lock() = Some(rt.handle().clone());
cvar.notify_all();
}
let stop_notifier = Arc::new(tokio::sync::Notify::new());
let stop_notifier_clone = stop_notifier.clone();
@@ -263,6 +289,43 @@ impl EasyTierLauncher {
}
}
}
#[cfg(feature = "ffi-dataplane")]
pub fn get_data_plane(&self) -> Option<Arc<Socks5Server>> {
self.data.data_plane.borrow().clone()
}
/// Waits up to `deadline` for the data-plane server to be published.
#[cfg(feature = "ffi-dataplane")]
pub async fn wait_data_plane(
&self,
deadline: tokio::time::Instant,
) -> Option<Arc<Socks5Server>> {
let mut rx = self.data.data_plane.subscribe();
loop {
if let Some(server) = rx.borrow_and_update().clone() {
return Some(server);
}
if tokio::time::timeout_at(deadline, rx.changed())
.await
.is_err()
{
return None;
}
}
}
/// Blocks up to `timeout` for the runtime handle to be published.
#[cfg(feature = "ffi-dataplane")]
pub fn wait_runtime_handle(
&self,
timeout: std::time::Duration,
) -> Option<tokio::runtime::Handle> {
let (lock, cvar) = &self.data.runtime_handle;
let mut guard = lock.lock();
cvar.wait_while_for(&mut guard, |h| h.is_none(), timeout);
guard.clone()
}
}
impl Default for EasyTierLauncher {
@@ -437,6 +500,10 @@ impl NetworkInstance {
&self.config_file_control
}
pub fn get_config(&self) -> TomlConfigLoader {
self.config.clone()
}
pub fn get_network_config_source(&self) -> ConfigSource {
self.config.get_network_config_source()
}
@@ -454,6 +521,75 @@ impl NetworkInstance {
.as_ref()
.and_then(|launcher| launcher.get_api_service())
}
/// Waits up to `timeout` for the data-plane server to come up, returning it
/// together with the deadline so the caller can spend the remaining budget
/// on the actual operation.
#[cfg(feature = "ffi-dataplane")]
async fn wait_data_plane(
&self,
timeout: std::time::Duration,
) -> anyhow::Result<(Arc<Socks5Server>, tokio::time::Instant)> {
let deadline = tokio::time::Instant::now() + timeout;
let launcher = self
.launcher
.as_ref()
.ok_or_else(|| anyhow::anyhow!("data plane is not ready"))?;
let server = launcher
.wait_data_plane(deadline)
.await
.ok_or_else(|| anyhow::anyhow!("data plane is not ready"))?;
Ok((server, deadline))
}
#[cfg(feature = "ffi-dataplane")]
pub async fn data_plane_tcp_connect(
&self,
dst_addr: SocketAddr,
timeout: std::time::Duration,
) -> anyhow::Result<DataPlaneTcpStream> {
let (server, deadline) = self.wait_data_plane(timeout).await?;
server
.data_plane_tcp_connect(dst_addr, deadline - tokio::time::Instant::now())
.await
.map_err(Into::into)
}
#[cfg(feature = "ffi-dataplane")]
pub async fn data_plane_tcp_bind(
&self,
local_port: u16,
timeout: std::time::Duration,
) -> anyhow::Result<DataPlaneTcpListener> {
let (server, deadline) = self.wait_data_plane(timeout).await?;
server
.data_plane_tcp_bind(local_port, deadline - tokio::time::Instant::now())
.await
.map_err(Into::into)
}
#[cfg(feature = "ffi-dataplane")]
pub async fn data_plane_udp_bind(
&self,
local_port: u16,
timeout: std::time::Duration,
) -> anyhow::Result<DataPlaneUdpSocket> {
let (server, deadline) = self.wait_data_plane(timeout).await?;
server
.data_plane_udp_bind(local_port, deadline - tokio::time::Instant::now())
.await
.map_err(Into::into)
}
#[cfg(feature = "ffi-dataplane")]
pub fn wait_runtime_handle(
&self,
timeout: std::time::Duration,
) -> Option<tokio::runtime::Handle> {
self.launcher
.as_ref()
.and_then(|launcher| launcher.wait_runtime_handle(timeout))
}
}
pub fn add_proxy_network_to_config(
@@ -768,6 +904,10 @@ impl NetworkConfig {
flags.bind_device = bind_device;
}
if self.socket_mark.is_some() {
flags.socket_mark = self.socket_mark;
}
if let Some(no_tun) = self.no_tun {
flags.no_tun = no_tun;
}
@@ -988,6 +1128,7 @@ impl NetworkConfig {
result.p2p_only = Some(flags.p2p_only);
result.lazy_p2p = Some(flags.lazy_p2p);
result.bind_device = Some(flags.bind_device);
result.socket_mark = flags.socket_mark;
result.no_tun = Some(flags.no_tun);
result.enable_exit_node = Some(flags.enable_exit_node);
result.relay_all_peer_rpc = Some(flags.relay_all_peer_rpc);
+81
View File
@@ -127,6 +127,8 @@ impl CredentialManager {
credential_id: Option<String>,
reusable: bool,
) -> (String, String) {
self.remove_expired_credentials();
let mut credentials = self.credentials.lock().unwrap();
let id = if let Some(id) = credential_id
.map(|x| x.trim().to_string())
@@ -194,6 +196,25 @@ impl CredentialManager {
removed
}
pub fn remove_expired_credentials(&self) -> bool {
self.remove_expired_credentials_at(current_unix_timestamp())
}
fn remove_expired_credentials_at(&self, now: i64) -> bool {
let removed = {
let mut credentials = self.credentials.lock().unwrap();
let before = credentials.len();
credentials.retain(|_, entry| entry.is_active_at(now));
before != credentials.len()
};
if removed {
self.save_to_disk();
}
removed
}
pub fn get_trusted_pubkeys(&self, network_secret: &str) -> Vec<TrustedCredentialPubkeyProof> {
let now = current_unix_timestamp();
@@ -496,6 +517,35 @@ mod tests {
assert_eq!(list.len(), 1);
}
#[test]
fn test_remove_expired_credentials_removes_and_persists() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("creds.json");
let mgr = CredentialManager::new(Some(path.clone()));
mgr.generate_credential_with_id(
vec!["active".to_string()],
false,
vec![],
Duration::from_secs(3600),
Some("active-id".to_string()),
);
mgr.generate_credential_with_id(
vec!["expired".to_string()],
false,
vec![],
Duration::from_secs(0),
Some("expired-id".to_string()),
);
assert!(mgr.remove_expired_credentials());
assert_eq!(mgr.list_credentials().len(), 1);
let reloaded = CredentialManager::new(Some(path));
let list = reloaded.list_credentials();
assert_eq!(list.len(), 1);
assert_eq!(list[0].credential_id, "active-id");
}
#[test]
fn test_generate_with_specified_id_reuses_existing_result() {
let mgr = CredentialManager::new(None);
@@ -528,6 +578,37 @@ mod tests {
assert_eq!(list[0].reusable, Some(true));
}
#[test]
fn test_generate_with_specified_id_replaces_expired_existing_result() {
let mgr = CredentialManager::new(None);
let fixed_id = "fixed-credential-id".to_string();
let (id1, secret1) = mgr.generate_credential_with_id(
vec!["expired".to_string()],
false,
vec![],
Duration::from_secs(0),
Some(fixed_id.clone()),
);
let (id2, secret2) = mgr.generate_credential_with_id(
vec!["fresh".to_string()],
true,
vec!["10.0.0.0/24".to_string()],
Duration::from_secs(3600),
Some(fixed_id.clone()),
);
assert_eq!(id1, fixed_id);
assert_eq!(id2, fixed_id);
assert_ne!(secret1, secret2);
let list = mgr.list_credentials();
assert_eq!(list.len(), 1);
assert_eq!(list[0].credential_id, fixed_id);
assert_eq!(list[0].groups, vec!["fresh".to_string()]);
assert!(list[0].allow_relay);
assert_eq!(list[0].allowed_proxy_cidrs, vec!["10.0.0.0/24".to_string()]);
}
#[test]
fn test_generate_non_reusable_credential() {
let mgr = CredentialManager::new(None);
@@ -14,7 +14,7 @@ use std::{
};
use dashmap::{DashMap, DashSet};
use guarden::defer;
use guarden::{Guard, defer};
use tokio::{
sync::{
Mutex,
@@ -284,6 +284,10 @@ impl ForeignNetworkEntry {
let mut flags = config.get_flags();
flags.disable_relay_kcp = !global_ctx.get_flags().enable_relay_foreign_network_kcp;
flags.disable_relay_quic = !global_ctx.get_flags().enable_relay_foreign_network_quic;
// socket_mark is a host-wide socket option: propagate from parent so
// outbound sockets the foreign-network entry initiates inherit the same
// mark as the rest of the node.
flags.socket_mark = global_ctx.get_flags().socket_mark;
config.set_flags(flags);
config.set_mapped_listeners(Some(global_ctx.config.get_mapped_listeners()));
+4 -5
View File
@@ -33,7 +33,6 @@ use super::{
peer_session::{PeerSession, PeerSessionAction},
traffic_metrics::AggregateTrafficMetrics,
};
use crate::utils::BoxExt;
use crate::{
common::{
PeerId,
@@ -380,9 +379,9 @@ impl PeerConn {
session_filter,
noise_handshake_result: None,
tunnel: Arc::new(Mutex::new(
guard!([mut mpsc_tunnel] mpsc_tunnel.close()).boxed(),
)),
tunnel: Arc::new(Mutex::new(Box::new(
guard!([mut mpsc_tunnel] mpsc_tunnel.close()),
))),
sink,
recv: Mutex::new(Some(recv)),
tunnel_info,
@@ -654,7 +653,7 @@ impl PeerConn {
.and_then(|p| p.peer_public_key)
}
async fn send_noise_msg<Msg: prost::Message>(
async fn send_noise_msg<Msg: prost::Message + Debug>(
&self,
pb: Msg,
packet_type: PacketType,
+267 -4
View File
@@ -503,6 +503,33 @@ impl PeerManager {
});
}
async fn close_untrusted_credential_peers(peer_map: &Arc<PeerMap>, global_ctx: &ArcGlobalCtx) {
let network_name = global_ctx.get_network_name();
for peer_id in peer_map.list_peers() {
if !matches!(
peer_map.get_peer_identity_type(peer_id),
Some(PeerIdentityType::Credential)
) {
continue;
}
let Some(peer) = peer_map.get_peer_by_id(peer_id) else {
continue;
};
let Some(pubkey) = peer.get_peer_public_key() else {
continue;
};
if global_ctx.is_pubkey_trusted(&pubkey, &network_name) {
continue;
}
tracing::warn!(?peer_id, "closing untrusted credential peer");
if let Err(e) = peer_map.close_peer(peer_id).await {
tracing::warn!(?e, ?peer_id, "failed to close untrusted credential peer");
}
}
}
fn build_foreign_network_manager_accessor(
peer_map: &Arc<PeerMap>,
) -> Box<dyn GlobalForeignNetworkAccessor> {
@@ -1397,7 +1424,7 @@ impl PeerManager {
let f = OneForeignNetwork {
network_name: info.key.as_ref().unwrap().network_name.clone(),
peer_ids: route_info.foreign_peer_ids.clone(),
last_updated: format!("{}", route_info.last_update.unwrap()),
last_updated: serde_json::to_string(&route_info.last_update.unwrap()).unwrap(),
version: route_info.version,
};
@@ -1506,9 +1533,22 @@ impl PeerManager {
) -> Result<(), Error> {
let policy =
Self::get_next_hop_policy(msg.peer_manager_header().unwrap().is_latency_first());
let is_latency_first = msg.peer_manager_header().unwrap().is_latency_first();
let packet_type = msg.peer_manager_header().unwrap().packet_type;
let msg_len = msg.buf_len() as u64;
let send_result = if peers.has_peer(dst_peer_id) {
let latency_first_gateway = if is_latency_first {
peers
.get_gateway_peer_id(dst_peer_id, policy.clone())
.await
.filter(|gateway| *gateway != dst_peer_id)
} else {
None
};
let send_result = if let Some(gateway) = latency_first_gateway
&& (peers.has_peer(gateway) || foreign_network_client.has_next_hop(gateway))
{
relay_peer_map.send_msg(msg, dst_peer_id, policy).await
} else if peers.has_peer(dst_peer_id) {
peers.send_msg_directly(msg, dst_peer_id).await
} else if foreign_network_client.has_next_hop(dst_peer_id) {
foreign_network_client.send_msg(msg, dst_peer_id).await
@@ -1849,6 +1889,26 @@ impl PeerManager {
});
}
async fn run_credential_gc_routine(&self) {
let global_ctx = self.global_ctx.clone();
let peer_map = self.peers.clone();
self.tasks.lock().await.spawn(async move {
loop {
if global_ctx.get_network_identity().network_secret.is_some() {
if global_ctx
.get_credential_manager()
.remove_expired_credentials()
{
global_ctx.issue_event(GlobalCtxEvent::CredentialChanged);
}
Self::close_untrusted_credential_peers(&peer_map, &global_ctx).await;
}
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
}
});
}
async fn run_traffic_metrics_gc_routine(&self) {
let mut event_receiver = self.global_ctx.subscribe();
let traffic_metrics = self.traffic_metrics.clone();
@@ -1897,6 +1957,7 @@ impl PeerManager {
self.run_relay_session_gc_routine().await;
self.run_recent_traffic_gc_routine().await;
self.run_peer_session_gc_routine().await;
self.run_credential_gc_routine().await;
self.run_traffic_metrics_gc_routine().await;
self.run_foriegn_network().await;
@@ -2135,7 +2196,9 @@ impl PeerManager {
#[cfg(test)]
mod tests {
use base64::Engine;
use std::{
collections::HashMap,
fmt::Debug,
sync::Arc,
time::{Duration, Instant},
@@ -2143,6 +2206,7 @@ mod tests {
use crate::{
common::{
PeerId,
config::Flags,
global_ctx::{NetworkIdentity, tests::get_mock_global_ctx},
stats_manager::{LabelSet, LabelType, MetricName},
@@ -2157,14 +2221,14 @@ mod tests {
peer_conn::tests::set_secure_mode_cfg,
peer_manager::RouteAlgoType,
peer_rpc::tests::register_service,
route_trait::NextHopPolicy,
route_trait::{NextHopPolicy, RouteCostCalculatorInterface},
tests::{
connect_peer_manager, create_mock_peer_manager_with_name, wait_route_appear,
wait_route_appear_with_cost,
},
},
proto::{
common::{CompressionAlgoPb, NatType},
common::{CompressionAlgoPb, NatType, SecureModeConfig},
peer_rpc::SecureAuthLevel,
},
tunnel::{
@@ -2201,6 +2265,16 @@ mod tests {
))
}
struct TestCostCalculator {
costs: HashMap<(PeerId, PeerId), i32>,
}
impl RouteCostCalculatorInterface for TestCostCalculator {
fn calculate_cost(&self, src: PeerId, dst: PeerId) -> i32 {
*self.costs.get(&(src, dst)).unwrap_or(&1)
}
}
#[test]
fn recent_traffic_fanout_policy_only_marks_single_peer() {
assert!(PeerManager::should_mark_recent_traffic_for_fanout(0));
@@ -2608,6 +2682,109 @@ mod tests {
.await;
}
#[tokio::test]
async fn send_msg_internal_uses_latency_first_gateway_for_direct_peer() {
let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await;
let peer_mgr_b = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await;
let peer_mgr_c = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await;
connect_peer_manager(peer_mgr_a.clone(), peer_mgr_b.clone()).await;
connect_peer_manager(peer_mgr_b.clone(), peer_mgr_c.clone()).await;
connect_peer_manager(peer_mgr_a.clone(), peer_mgr_c.clone()).await;
wait_route_appear(peer_mgr_a.clone(), peer_mgr_b.clone())
.await
.unwrap();
wait_route_appear(peer_mgr_b.clone(), peer_mgr_c.clone())
.await
.unwrap();
wait_route_appear(peer_mgr_a.clone(), peer_mgr_c.clone())
.await
.unwrap();
peer_mgr_a
.get_route()
.set_route_cost_fn(Box::new(TestCostCalculator {
costs: HashMap::from([
((peer_mgr_a.my_peer_id(), peer_mgr_c.my_peer_id()), 100),
((peer_mgr_a.my_peer_id(), peer_mgr_b.my_peer_id()), 1),
((peer_mgr_b.my_peer_id(), peer_mgr_c.my_peer_id()), 1),
]),
}))
.await;
wait_for_condition(
|| {
let peer_mgr_a = peer_mgr_a.clone();
let peer_mgr_b = peer_mgr_b.clone();
let peer_mgr_c = peer_mgr_c.clone();
async move {
peer_mgr_a
.get_route()
.get_next_hop_with_policy(peer_mgr_c.my_peer_id(), NextHopPolicy::LeastCost)
.await
== Some(peer_mgr_b.my_peer_id())
}
},
Duration::from_secs(5),
)
.await;
let b_network_labels = network_labels(&peer_mgr_b);
let forwarded_bytes_before = metric_value(
&peer_mgr_b,
MetricName::TrafficBytesForwarded,
&b_network_labels,
);
let forwarded_packets_before = metric_value(
&peer_mgr_b,
MetricName::TrafficPacketsForwarded,
&b_network_labels,
);
let mut pkt = ZCPacket::new_with_payload(b"latency-first");
pkt.fill_peer_manager_hdr(
peer_mgr_a.my_peer_id(),
peer_mgr_c.my_peer_id(),
PacketType::Data as u8,
);
pkt.mut_peer_manager_header()
.unwrap()
.set_latency_first(true);
let pkt_len = pkt.buf_len() as u64;
PeerManager::send_msg_internal(
&peer_mgr_a.peers,
&peer_mgr_a.foreign_network_client,
&peer_mgr_a.relay_peer_map,
Some(&peer_mgr_a.traffic_metrics),
pkt,
peer_mgr_c.my_peer_id(),
)
.await
.unwrap();
wait_for_condition(
|| {
let peer_mgr_b = peer_mgr_b.clone();
let b_network_labels = b_network_labels.clone();
async move {
metric_value(
&peer_mgr_b,
MetricName::TrafficBytesForwarded,
&b_network_labels,
) >= forwarded_bytes_before + pkt_len
&& metric_value(
&peer_mgr_b,
MetricName::TrafficPacketsForwarded,
&b_network_labels,
) > forwarded_packets_before
}
},
Duration::from_secs(5),
)
.await;
}
#[tokio::test]
async fn send_msg_internal_records_control_metrics_for_direct_peer() {
let peer_mgr_a = create_mock_peer_manager_with_mock_stun(NatType::Unknown).await;
@@ -3406,6 +3583,92 @@ mod tests {
// a is client, b is server
}
#[tokio::test]
async fn expired_credential_peer_conn_is_closed_without_ospf() {
let (admin_ch, _admin_rx) = create_packet_recv_chan();
let admin_ctx = get_mock_global_ctx();
admin_ctx.config.set_network_identity(NetworkIdentity::new(
"net1".to_string(),
"secret".to_string(),
));
set_secure_mode_cfg(&admin_ctx, true);
let admin = Arc::new(PeerManager::new(
RouteAlgoType::None,
admin_ctx.clone(),
admin_ch,
));
admin.run().await.unwrap();
let (_cred_id, cred_secret) = admin_ctx.get_credential_manager().generate_credential(
vec![],
false,
vec![],
Duration::from_secs(1),
);
let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD
.decode(&cred_secret)
.unwrap()
.try_into()
.unwrap();
let private = x25519_dalek::StaticSecret::from(privkey_bytes);
let public = x25519_dalek::PublicKey::from(&private);
let (credential_ch, _credential_rx) = create_packet_recv_chan();
let credential_ctx = get_mock_global_ctx();
credential_ctx
.config
.set_network_identity(NetworkIdentity::new_credential("net1".to_string()));
credential_ctx
.config
.set_secure_mode(Some(SecureModeConfig {
enabled: true,
local_private_key: Some(
base64::engine::general_purpose::STANDARD.encode(private.as_bytes()),
),
local_public_key: Some(
base64::engine::general_purpose::STANDARD.encode(public.as_bytes()),
),
}));
let credential = Arc::new(PeerManager::new(
RouteAlgoType::None,
credential_ctx,
credential_ch,
));
credential.run().await.unwrap();
let credential_peer_id = credential.my_peer_id();
connect_peer_manager(credential.clone(), admin.clone()).await;
wait_for_condition(
|| {
let admin = admin.clone();
async move {
admin
.get_peer_map()
.list_peer_conns(credential_peer_id)
.await
.is_some_and(|conns| !conns.is_empty())
}
},
Duration::from_secs(5),
)
.await;
wait_for_condition(
|| {
let admin = admin.clone();
async move {
admin
.get_peer_map()
.list_peer_conns(credential_peer_id)
.await
.is_none_or(|conns| conns.is_empty())
}
},
Duration::from_secs(5),
)
.await;
}
#[tokio::test]
async fn close_conn_in_foreign_network_client() {
let peer_mgr_server = create_mock_peer_manager_with_name("server".to_string()).await;
+538 -92
View File
@@ -73,7 +73,9 @@ use super::{
},
};
use crate::proto::common::TimestampExt;
use atomic_shim::AtomicU64;
use prost_wkt_types::Timestamp;
static SERVICE_ID: u32 = 7;
static UPDATE_PEER_INFO_PERIOD: Duration = Duration::from_secs(3600);
@@ -439,6 +441,9 @@ struct SyncedRouteInfo {
// Tracks the currently accepted peer for non-reusable credentials.
// Maps credential pubkey bytes -> peer_id.
non_reusable_credential_owners: DashMap<Vec<u8>, PeerId>,
// Duplicate non-reusable credential peers are kept for OSPF sync and topology
// reachability, but excluded from forwarding until owner election selects them.
suppressed_non_reusable_credential_peers: DashMap<PeerId, ()>,
version: AtomicVersion,
}
@@ -658,6 +663,36 @@ impl SyncedRouteInfo {
}
}
fn replace_suppressed_non_reusable_credential_peers(
&self,
suppressed_peers: BTreeSet<PeerId>,
) -> bool {
let current: BTreeSet<_> = self
.suppressed_non_reusable_credential_peers
.iter()
.map(|entry| *entry.key())
.collect();
if current == suppressed_peers {
return false;
}
self.suppressed_non_reusable_credential_peers
.retain(|peer_id, _| suppressed_peers.contains(peer_id));
for peer_id in suppressed_peers {
self.suppressed_non_reusable_credential_peers
.insert(peer_id, ());
}
self.version.inc();
true
}
fn is_route_suppressed(&self, peer_id: PeerId) -> bool {
self.suppressed_non_reusable_credential_peers
.contains_key(&peer_id)
}
fn update_credential_groups(
&self,
peer_infos: &OrderedHashMap<PeerId, RoutePeerInfo>,
@@ -757,7 +792,7 @@ impl SyncedRouteInfo {
if !guard.contains_key(peer_id) {
let mut peer_info = RoutePeerInfo::new();
let mut guard = RwLockUpgradableReadGuard::upgrade(guard);
peer_info.last_update = Some(SystemTime::now().into());
peer_info.last_update = Some(Timestamp::now());
guard.insert(*peer_id, peer_info);
need_inc_version = true;
} else {
@@ -867,7 +902,7 @@ impl SyncedRouteInfo {
let mut guard = self.peer_infos.write();
// time between peers may not be synchronized, so update last_update to local now.
// note only last_update with larger version will be updated to local saved peer info.
route_info.last_update = Some(SystemTime::now().into());
route_info.last_update = Some(Timestamp::now());
if guard
.get_mut(&route_info.peer_id)
.is_none_or(|old| route_info.version > old.version)
@@ -962,7 +997,7 @@ impl SyncedRouteInfo {
continue;
};
entry.last_update = Some(SystemTime::now().into());
entry.last_update = Some(Timestamp::now());
self.foreign_network
.entry(key.clone())
@@ -1010,7 +1045,7 @@ impl SyncedRouteInfo {
};
guard.with_upgraded(|peer_infos| {
new.last_update = Some(SystemTime::now().into());
new.last_update = Some(Timestamp::now());
new.version = new_version;
peer_infos.insert(my_peer_id, new)
});
@@ -1086,7 +1121,7 @@ impl SyncedRouteInfo {
foreign_networks.remove(key).unwrap();
} else if !item.foreign_peer_ids.is_empty() {
item.foreign_peer_ids.clear();
item.last_update = Some(SystemTime::now().into());
item.last_update = Some(Timestamp::now());
item.version = std::cmp::max(item.version + 1, now_version);
updated = true;
}
@@ -1231,11 +1266,13 @@ impl SyncedRouteInfo {
where
F: FnMut(PeerId) -> bool,
{
self.verify_and_update_credential_trusts_with_active_peers_protecting(
network_secret,
is_peer_active,
None,
)
let (untrusted_peers, global_trusted_keys, _) = self
.verify_and_update_credential_trusts_with_active_peers_protecting(
network_secret,
is_peer_active,
None,
);
(untrusted_peers, global_trusted_keys)
}
fn verify_and_update_credential_trusts_with_active_peers_protecting<F>(
@@ -1246,6 +1283,7 @@ impl SyncedRouteInfo {
) -> (
Vec<PeerId>,
HashMap<Vec<u8>, crate::common::global_ctx::TrustedKeyMetadata>,
bool,
)
where
F: FnMut(PeerId) -> bool,
@@ -1259,14 +1297,18 @@ impl SyncedRouteInfo {
let (all_trusted, global_trusted_keys) =
self.collect_trusted_credentials(&peer_infos, network_secret, now);
let prev_trusted = self.replace_trusted_credential_pubkeys(&all_trusted);
let (active_non_reusable_owners, duplicate_untrusted_peers) =
let (active_non_reusable_owners, mut duplicate_untrusted_peers) =
self.collect_non_reusable_credential_owners(&peer_infos, &all_trusted, is_peer_active);
if let Some(protected_peer_id) = protected_peer_id {
duplicate_untrusted_peers.remove(&protected_peer_id);
}
self.replace_non_reusable_credential_owners(active_non_reusable_owners);
let suppressed_changed =
self.replace_suppressed_non_reusable_credential_peers(duplicate_untrusted_peers);
self.update_credential_groups(&peer_infos, &all_trusted);
let mut untrusted_peers =
Self::collect_revoked_credential_peers(&peer_infos, &prev_trusted, &all_trusted);
untrusted_peers.extend(duplicate_untrusted_peers);
if let Some(protected_peer_id) = protected_peer_id {
untrusted_peers.remove(&protected_peer_id);
}
@@ -1280,7 +1322,11 @@ impl SyncedRouteInfo {
self.remove_peers(untrusted_peers.iter().copied());
}
(untrusted_peers.into_iter().collect(), global_trusted_keys)
(
untrusted_peers.into_iter().collect(),
global_trusted_keys,
suppressed_changed,
)
}
fn is_admin_peer(&self, info: &RoutePeerInfo) -> bool {
@@ -1325,6 +1371,7 @@ type NextHopMap = DashMap<PeerId, NextHopInfo>;
struct RouteTable {
peer_infos: DashMap<PeerId, RoutePeerInfo>,
next_hop_map: NextHopMap,
suppressed_peer_ids: DashMap<PeerId, ()>,
ipv4_peer_id_map: DashMap<Ipv4Addr, PeerIdVersion>,
ipv6_peer_id_map: DashMap<Ipv6Addr, PeerIdVersion>,
cidr_peer_id_map: ArcSwap<PrefixMap<Ipv4Cidr, PeerIdVersion>>,
@@ -1337,6 +1384,7 @@ impl RouteTable {
RouteTable {
peer_infos: DashMap::new(),
next_hop_map: DashMap::new(),
suppressed_peer_ids: DashMap::new(),
ipv4_peer_id_map: DashMap::new(),
ipv6_peer_id_map: DashMap::new(),
cidr_peer_id_map: ArcSwap::new(Arc::new(PrefixMap::new())),
@@ -1346,6 +1394,13 @@ impl RouteTable {
}
fn get_next_hop(&self, dst_peer_id: PeerId) -> Option<NextHopInfo> {
if self.suppressed_peer_ids.contains_key(&dst_peer_id) {
return None;
}
self.get_topology_next_hop(dst_peer_id)
}
fn get_topology_next_hop(&self, dst_peer_id: PeerId) -> Option<NextHopInfo> {
let cur_version = self.next_hop_map_version.get();
self.next_hop_map.get(&dst_peer_id).and_then(|x| {
if x.version >= cur_version {
@@ -1360,6 +1415,18 @@ impl RouteTable {
self.get_next_hop(peer_id).is_some()
}
fn topology_peer_reachable(&self, peer_id: PeerId) -> bool {
self.get_topology_next_hop(peer_id).is_some()
}
fn sync_suppressed_peer_ids(&self, synced_info: &SyncedRouteInfo) {
self.suppressed_peer_ids
.retain(|peer_id, _| synced_info.is_route_suppressed(*peer_id));
for entry in synced_info.suppressed_non_reusable_credential_peers.iter() {
self.suppressed_peer_ids.insert(*entry.key(), ());
}
}
fn get_udp_nat_type(&self, peer_id: PeerId) -> Option<NatType> {
self.peer_infos
.get(&peer_id)
@@ -1396,21 +1463,24 @@ impl RouteTable {
}
for item in peer_id_to_node_index.iter() {
let src_peer_id = item.key();
let src_peer_id = *item.key();
if src_peer_id != my_peer_id && synced_info.is_route_suppressed(src_peer_id) {
continue;
}
let src_node_idx = item.value();
let connected_peers: BTreeSet<_> = synced_info
.get_connected_peers(*src_peer_id)
.get_connected_peers(src_peer_id)
.unwrap_or_default();
// if avoid relay, just set all outgoing edges to a large value: AVOID_RELAY_COST.
let peer_avoid_relay_data = synced_info.get_avoid_relay_data(*src_peer_id);
let peer_avoid_relay_data = synced_info.get_avoid_relay_data(src_peer_id);
for dst_peer_id in connected_peers.iter() {
let Some(dst_node_idx) = peer_id_to_node_index.get(dst_peer_id) else {
continue;
};
let mut cost = cost_calc.calculate_cost(*src_peer_id, *dst_peer_id) as usize;
let mut cost = cost_calc.calculate_cost(src_peer_id, *dst_peer_id) as usize;
if peer_avoid_relay_data {
cost += AVOID_RELAY_COST;
}
@@ -1429,20 +1499,21 @@ impl RouteTable {
v.version >= cur_version
});
self.peer_infos.retain(|k, _| {
// remove peer info for peers we cannot reach.
self.next_hop_map.contains_key(k)
// remove peer info for peers we cannot forward to.
self.peer_reachable(*k)
});
self.ipv4_peer_id_map.retain(|_, v| {
// remove ipv4 map for peers we cannot reach.
self.next_hop_map.contains_key(&v.peer_id)
// remove ipv4 map for peers we cannot forward to.
self.peer_reachable(v.peer_id)
});
self.ipv6_peer_id_map.retain(|_, v| {
// remove ipv6 map for peers we cannot reach.
self.next_hop_map.contains_key(&v.peer_id)
// remove ipv6 map for peers we cannot forward to.
self.peer_reachable(v.peer_id)
});
shrink_dashmap(&self.peer_infos, None);
shrink_dashmap(&self.next_hop_map, None);
shrink_dashmap(&self.suppressed_peer_ids, None);
shrink_dashmap(&self.ipv4_peer_id_map, None);
shrink_dashmap(&self.ipv6_peer_id_map, None);
}
@@ -1543,6 +1614,7 @@ impl RouteTable {
cost_calc: &T,
) {
let version = synced_info.version.get();
self.sync_suppressed_peer_ids(synced_info);
let local_proxy_cidrs = synced_info
.peer_infos
@@ -1592,6 +1664,10 @@ impl RouteTable {
}
let peer_id = item.key();
if !self.peer_reachable(*peer_id) {
continue;
}
let Some(info) = synced_info.peer_infos.read().get(peer_id).cloned() else {
continue;
};
@@ -1715,6 +1791,7 @@ impl RouteTable {
cidrs_v6 = ?self.cidr_v6_peer_id_map.load(),
"update peer cidr map"
);
self.clean_expired_route_info();
}
fn get_peer_id_for_proxy(&self, ip: &IpAddr) -> Option<PeerId> {
@@ -2145,6 +2222,7 @@ impl PeerRouteServiceImpl {
group_trust_map_cache: DashMap::new(),
trusted_credential_pubkeys: DashMap::new(),
non_reusable_credential_owners: DashMap::new(),
suppressed_non_reusable_credential_peers: DashMap::new(),
version: AtomicVersion::new(),
},
public_ipv6_service: std::sync::Mutex::new(Weak::new()),
@@ -2168,10 +2246,9 @@ impl PeerRouteServiceImpl {
ni.network_secret_digest.map(|d| d.to_vec())
}
#[cfg(test)]
fn is_active_non_reusable_credential_peer(&self, peer_id: PeerId) -> bool {
peer_id == self.my_peer_id
|| self.sessions.contains_key(&peer_id)
|| self.route_table.peer_reachable(peer_id)
peer_id == self.my_peer_id || self.route_table.topology_peer_reachable(peer_id)
}
fn is_credential_node(&self) -> bool {
@@ -2486,7 +2563,7 @@ impl PeerRouteServiceImpl {
};
for item in self.synced_route_info.conn_map.read().iter() {
let src_peer_id = *item.0;
if !self.route_table.peer_reachable(src_peer_id) {
if !self.route_table.topology_peer_reachable(src_peer_id) {
continue;
}
add_to_all_peer_ids(src_peer_id, item.1.version.get());
@@ -2542,7 +2619,7 @@ impl PeerRouteServiceImpl {
for (peer_id, peer_info) in peer_infos.iter().rev() {
// stop iter if last_update of peer info is older than session.last_sync_succ_timestamp
if let Some(last_update) = peer_info.last_update {
let last_update = TryInto::<SystemTime>::try_into(last_update).unwrap();
let last_update = SystemTime::try_from(last_update).unwrap();
if last_sync_succ_timestamp.is_some_and(|t| last_update < t) {
break;
}
@@ -2553,7 +2630,7 @@ impl PeerRouteServiceImpl {
}
// do not send unreachable peer info to dst peer.
if !self.route_table.peer_reachable(*peer_id) {
if !self.route_table.topology_peer_reachable(*peer_id) {
unreachable_peers_for_peer_info.insert(*peer_id, peer_info.version);
continue;
}
@@ -2571,7 +2648,7 @@ impl PeerRouteServiceImpl {
return false;
};
if self.route_table.peer_reachable(*peer_id) {
if self.route_table.topology_peer_reachable(*peer_id) {
route_infos.push(peer_info.clone());
}
@@ -2621,7 +2698,7 @@ impl PeerRouteServiceImpl {
continue;
}
if !self.route_table.peer_reachable(*peer_id) {
if !self.route_table.topology_peer_reachable(*peer_id) {
unreachable_peers_for_conn_info.insert(*peer_id, conn_info.version.get());
continue;
}
@@ -2639,7 +2716,7 @@ impl PeerRouteServiceImpl {
return false;
};
if self.route_table.peer_reachable(*peer_id) {
if self.route_table.topology_peer_reachable(*peer_id) {
add_to_conn_peer_list(*peer_id, conn_info);
}
@@ -2697,7 +2774,7 @@ impl PeerRouteServiceImpl {
let my_conn_info_updated = self.update_my_conn_info().await;
let my_foreign_network_updated = self.update_my_foreign_network().await;
let mut untrusted_changed = false;
if my_peer_info_updated {
if my_peer_info_updated || my_conn_info_updated {
untrusted_changed = self.refresh_credential_trusts_and_disconnect().await;
}
@@ -2755,7 +2832,7 @@ impl PeerRouteServiceImpl {
fn refresh_credential_trusts(&self) -> Vec<PeerId> {
let network_identity = self.global_ctx.get_network_identity();
let (untrusted, global_trusted_keys) = self
let (untrusted, global_trusted_keys, _) = self
.synced_route_info
.verify_and_update_credential_trusts_with_active_peers_protecting(
network_identity.network_secret.as_deref(),
@@ -2775,17 +2852,19 @@ impl PeerRouteServiceImpl {
// route table from the latest synced peer/conn state before checking active peers.
self.update_route_table_and_cached_local_conn_bitmap();
let (untrusted, global_trusted_keys) = self
let (untrusted, global_trusted_keys, suppressed_changed) = self
.synced_route_info
.verify_and_update_credential_trusts_with_active_peers_protecting(
network_identity.network_secret.as_deref(),
|peer_id| self.is_active_non_reusable_credential_peer(peer_id),
|peer_id| {
peer_id == self.my_peer_id || self.route_table.topology_peer_reachable(peer_id)
},
Some(self.my_peer_id),
);
self.global_ctx
.update_trusted_keys(global_trusted_keys, &network_identity.network_name);
if !untrusted.is_empty() {
if !untrusted.is_empty() || suppressed_changed {
self.update_route_table_and_cached_local_conn_bitmap();
}
untrusted
@@ -2864,7 +2943,7 @@ impl PeerRouteServiceImpl {
if let Ok(d) = now.duration_since(peer_info.last_update.unwrap().try_into().unwrap())
&& (d > REMOVE_DEAD_PEER_INFO_AFTER
|| (d > REMOVE_UNREACHABLE_PEER_INFO_AFTER
&& !self.route_table.peer_reachable(*peer_id)))
&& !self.route_table.topology_peer_reachable(*peer_id)))
{
to_remove.push(*peer_id);
}
@@ -4158,7 +4237,9 @@ mod tests {
use dashmap::DashMap;
use parking_lot::Mutex;
use prefix_trie::PrefixMap;
use prost::Message;
use prost_reflect::{DynamicMessage, ReflectMessage};
use prost_wkt_types::Timestamp;
use std::net::IpAddr;
use std::{
collections::{BTreeSet, HashMap},
@@ -4169,7 +4250,10 @@ mod tests {
time::{Duration, SystemTime},
};
use super::{NextHopInfo, PeerRoute, REMOVE_DEAD_PEER_INFO_AFTER, RouteConnInfo};
use super::{
NextHopInfo, PeerRoute, REMOVE_DEAD_PEER_INFO_AFTER, RouteConnInfo, SyncRouteSession,
};
use crate::proto::common::TimestampExt;
use crate::{
common::{
PeerId,
@@ -4198,7 +4282,9 @@ mod tests {
},
tunnel::common::tests::wait_for_condition,
};
use prost::Message;
use base64::Engine as _;
use base64::prelude::BASE64_STANDARD;
struct AuthOnlyInterface {
my_peer_id: PeerId,
identity_type: DashMap<PeerId, PeerIdentityType>,
@@ -4433,6 +4519,31 @@ mod tests {
peer_info
}
fn make_admin_route_peer_info(
peer_id: PeerId,
credential_key: &[u8],
network_secret: &str,
now: i64,
) -> RoutePeerInfo {
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = peer_id;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.to_vec(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
network_secret,
)];
admin_info
}
fn make_route_conn_info<I>(connected_peers: I, last_update: SystemTime) -> RouteConnInfo
where
I: IntoIterator<Item = PeerId>,
@@ -4882,22 +4993,7 @@ mod tests {
let credential_key = vec![7; 32];
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = 30;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.clone(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
network_secret,
)];
let admin_info = make_admin_route_peer_info(30, &credential_key, network_secret, now);
let mut original_peer = RoutePeerInfo::new();
original_peer.peer_id = 41;
@@ -4948,9 +5044,9 @@ mod tests {
let (second_untrusted, _) = service_impl
.synced_route_info
.verify_and_update_credential_trusts(Some(network_secret));
assert_eq!(second_untrusted, vec![41]);
assert!(second_untrusted.is_empty());
assert!(
!service_impl
service_impl
.synced_route_info
.peer_infos
.read()
@@ -4971,6 +5067,8 @@ mod tests {
.map(|entry| *entry.value()),
Some(39)
);
assert!(service_impl.synced_route_info.is_route_suppressed(41));
assert!(!service_impl.synced_route_info.is_route_suppressed(39));
}
#[tokio::test]
@@ -4986,22 +5084,7 @@ mod tests {
let stale_peer_id = 41;
let replacement_peer_id = 39;
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = 30;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.clone(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
network_secret,
)];
let admin_info = make_admin_route_peer_info(30, &credential_key, network_secret, now);
let mut stale_peer = RoutePeerInfo::new();
stale_peer.peer_id = stale_peer_id;
@@ -5072,6 +5155,292 @@ mod tests {
.map(|entry| *entry.value()),
Some(replacement_peer_id)
);
assert!(
!service_impl
.synced_route_info
.is_route_suppressed(stale_peer_id)
);
assert!(
!service_impl
.synced_route_info
.is_route_suppressed(replacement_peer_id)
);
}
#[tokio::test]
async fn suppressed_non_reusable_credential_peer_stays_synced_and_can_be_reactivated() {
const NETWORK_SECRET: &str = "sec1";
const SELF_PEER_ID: PeerId = 1;
const ADMIN_PEER_ID: PeerId = 30;
const FIRST_PEER_ID: PeerId = 39;
const SECOND_PEER_ID: PeerId = 41;
let service_impl = PeerRouteServiceImpl::new(
SELF_PEER_ID,
get_mock_global_ctx_with_network(Some(NetworkIdentity::new(
"test-net".to_string(),
NETWORK_SECRET.to_string(),
))),
);
let now_unix = SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64;
let now = SystemTime::now();
let credential_key = vec![10; 32];
let mut self_info = RoutePeerInfo::new();
self_info.peer_id = SELF_PEER_ID;
self_info.version = 1;
let admin_info =
make_admin_route_peer_info(ADMIN_PEER_ID, &credential_key, NETWORK_SECRET, now_unix);
let mut first_peer = make_credential_route_peer_info(FIRST_PEER_ID, &credential_key);
first_peer.ipv4_addr = Some(std::net::Ipv4Addr::new(10, 144, 0, 39).into());
let mut second_peer = make_credential_route_peer_info(SECOND_PEER_ID, &credential_key);
second_peer.ipv4_addr = Some(std::net::Ipv4Addr::new(10, 144, 0, 41).into());
second_peer.proxy_cidrs.push("10.244.41.0/24".into());
{
let mut peer_infos = service_impl.synced_route_info.peer_infos.write();
peer_infos.insert(self_info.peer_id, self_info);
peer_infos.insert(admin_info.peer_id, admin_info);
peer_infos.insert(first_peer.peer_id, first_peer);
peer_infos.insert(second_peer.peer_id, second_peer);
}
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(SELF_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(
ADMIN_PEER_ID,
make_route_conn_info([SELF_PEER_ID, FIRST_PEER_ID, SECOND_PEER_ID], now),
);
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
}
service_impl.synced_route_info.version.set(1);
let first_untrusted = service_impl.refresh_credential_trusts_with_current_topology();
assert!(first_untrusted.is_empty());
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(FIRST_PEER_ID)
);
assert!(
service_impl
.synced_route_info
.peer_infos
.read()
.contains_key(&SECOND_PEER_ID)
);
assert!(
service_impl
.synced_route_info
.is_route_suppressed(SECOND_PEER_ID)
);
assert!(
service_impl
.route_table
.topology_peer_reachable(SECOND_PEER_ID)
);
assert!(service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID));
assert!(
service_impl
.route_table
.peer_infos
.contains_key(&FIRST_PEER_ID)
);
assert!(
!service_impl
.route_table
.peer_infos
.contains_key(&SECOND_PEER_ID)
);
assert_eq!(
service_impl
.route_table
.ipv4_peer_id_map
.get(&"10.144.0.41".parse().unwrap())
.map(|entry| entry.peer_id),
None
);
assert_eq!(
service_impl
.route_table
.get_peer_id_for_proxy(&"10.244.41.1".parse().unwrap()),
None
);
let sync_session = SyncRouteSession::new(SELF_PEER_ID, ADMIN_PEER_ID);
let sync_peer_ids: BTreeSet<_> = service_impl
.build_route_info(&sync_session)
.unwrap()
.into_iter()
.map(|info| info.peer_id)
.collect();
assert!(sync_peer_ids.contains(&SECOND_PEER_ID));
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(
ADMIN_PEER_ID,
make_route_conn_info([SELF_PEER_ID, SECOND_PEER_ID], now),
);
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([], now));
conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
}
service_impl.synced_route_info.version.inc();
let second_untrusted = service_impl.refresh_credential_trusts_with_current_topology();
assert!(second_untrusted.is_empty());
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(SECOND_PEER_ID)
);
assert!(
!service_impl
.synced_route_info
.is_route_suppressed(SECOND_PEER_ID)
);
assert!(
service_impl
.route_table
.topology_peer_reachable(SECOND_PEER_ID)
);
assert!(!service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(service_impl.route_table.peer_reachable(SECOND_PEER_ID));
assert!(
!service_impl
.route_table
.peer_infos
.contains_key(&FIRST_PEER_ID)
);
assert!(
service_impl
.route_table
.peer_infos
.contains_key(&SECOND_PEER_ID)
);
assert_eq!(
service_impl
.route_table
.ipv4_peer_id_map
.get(&"10.144.0.41".parse().unwrap())
.map(|entry| entry.peer_id),
Some(SECOND_PEER_ID)
);
assert_eq!(
service_impl
.route_table
.get_peer_id_for_proxy(&"10.244.41.1".parse().unwrap()),
Some(SECOND_PEER_ID)
);
}
#[tokio::test]
async fn suppressed_non_reusable_credential_peer_is_not_transit_next_hop() {
const NETWORK_SECRET: &str = "sec1";
const SELF_PEER_ID: PeerId = 1;
const ADMIN_PEER_ID: PeerId = 30;
const FIRST_PEER_ID: PeerId = 39;
const SECOND_PEER_ID: PeerId = 41;
const DOWNSTREAM_PEER_ID: PeerId = 50;
let service_impl = PeerRouteServiceImpl::new(
SELF_PEER_ID,
get_mock_global_ctx_with_network(Some(NetworkIdentity::new(
"test-net".to_string(),
NETWORK_SECRET.to_string(),
))),
);
let now_unix = SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64;
let now = SystemTime::now();
let credential_key = vec![10; 32];
let mut self_info = RoutePeerInfo::new();
self_info.peer_id = SELF_PEER_ID;
self_info.version = 1;
let admin_info =
make_admin_route_peer_info(ADMIN_PEER_ID, &credential_key, NETWORK_SECRET, now_unix);
let first_peer = make_credential_route_peer_info(FIRST_PEER_ID, &credential_key);
let second_peer = make_credential_route_peer_info(SECOND_PEER_ID, &credential_key);
let mut downstream_peer = RoutePeerInfo::new();
downstream_peer.peer_id = DOWNSTREAM_PEER_ID;
downstream_peer.version = 1;
{
let mut peer_infos = service_impl.synced_route_info.peer_infos.write();
peer_infos.insert(self_info.peer_id, self_info);
peer_infos.insert(admin_info.peer_id, admin_info);
peer_infos.insert(first_peer.peer_id, first_peer);
peer_infos.insert(second_peer.peer_id, second_peer);
peer_infos.insert(downstream_peer.peer_id, downstream_peer);
}
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(SELF_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(
ADMIN_PEER_ID,
make_route_conn_info([SELF_PEER_ID, FIRST_PEER_ID, SECOND_PEER_ID], now),
);
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(
SECOND_PEER_ID,
make_route_conn_info([ADMIN_PEER_ID, DOWNSTREAM_PEER_ID], now),
);
conn_map.insert(
DOWNSTREAM_PEER_ID,
make_route_conn_info([SECOND_PEER_ID], now),
);
}
service_impl.synced_route_info.version.set(1);
let untrusted = service_impl.refresh_credential_trusts_with_current_topology();
assert!(untrusted.is_empty());
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(FIRST_PEER_ID)
);
assert!(
service_impl
.synced_route_info
.is_route_suppressed(SECOND_PEER_ID)
);
assert!(
service_impl
.route_table
.topology_peer_reachable(SECOND_PEER_ID)
);
assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID));
assert!(!service_impl.route_table.peer_reachable(DOWNSTREAM_PEER_ID));
assert!(
service_impl
.route_table
.get_next_hop(DOWNSTREAM_PEER_ID)
.is_none()
);
assert!(
!service_impl
.route_table
.peer_infos
.contains_key(&DOWNSTREAM_PEER_ID)
);
}
#[tokio::test]
@@ -5101,7 +5470,7 @@ mod tests {
},
);
let (untrusted_peers, _) = service_impl
let (untrusted_peers, _, _) = service_impl
.synced_route_info
.verify_and_update_credential_trusts_with_active_peers_protecting(
None,
@@ -5152,22 +5521,8 @@ mod tests {
self_info.peer_id = SELF_PEER_ID;
self_info.version = 1;
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = admin_peer_id;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.clone(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
NETWORK_SECRET,
)];
let admin_info =
make_admin_route_peer_info(admin_peer_id, &credential_key, NETWORK_SECRET, now);
let stale_peer = make_credential_route_peer_info(stale_peer_id, &credential_key);
let replacement_peer =
@@ -5227,6 +5582,97 @@ mod tests {
);
}
#[tokio::test]
async fn update_my_infos_refreshes_non_reusable_owner_on_conn_change() {
const NETWORK_SECRET: &str = "sec1";
const ADMIN_PEER_ID: PeerId = 30;
const FIRST_PEER_ID: PeerId = 39;
const SECOND_PEER_ID: PeerId = 41;
let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new(
"test-net".to_string(),
NETWORK_SECRET.to_string(),
)));
let (_credential_id, credential_secret) = global_ctx
.get_credential_manager()
.generate_credential_with_options(
vec![],
false,
vec![],
Duration::from_secs(3600),
None,
false,
);
let credential_secret_bytes: [u8; 32] = BASE64_STANDARD
.decode(&credential_secret)
.unwrap()
.try_into()
.unwrap();
let credential_secret = x25519_dalek::StaticSecret::from(credential_secret_bytes);
let credential_key = x25519_dalek::PublicKey::from(&credential_secret)
.as_bytes()
.to_vec();
let service_impl = PeerRouteServiceImpl::new(ADMIN_PEER_ID, global_ctx);
let peers = Arc::new(Mutex::new(vec![FIRST_PEER_ID, SECOND_PEER_ID]));
let peer_identity_types = Arc::new(Mutex::new(HashMap::from([
(FIRST_PEER_ID, Some(PeerIdentityType::Credential)),
(SECOND_PEER_ID, Some(PeerIdentityType::Credential)),
])));
*service_impl.interface.lock().await = Some(Box::new(CountingInterface {
my_peer_id: ADMIN_PEER_ID,
peers: peers.clone(),
peer_identity_types,
list_peers_calls: Arc::new(AtomicU32::new(0)),
get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)),
}));
{
let mut peer_infos = service_impl.synced_route_info.peer_infos.write();
peer_infos.insert(
FIRST_PEER_ID,
make_credential_route_peer_info(FIRST_PEER_ID, &credential_key),
);
peer_infos.insert(
SECOND_PEER_ID,
make_credential_route_peer_info(SECOND_PEER_ID, &credential_key),
);
}
let now = SystemTime::now();
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
}
assert!(service_impl.update_my_infos().await);
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(FIRST_PEER_ID)
);
assert!(service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID));
*peers.lock() = vec![SECOND_PEER_ID];
service_impl.handle_global_ctx_event(&GlobalCtxEvent::PeerConnRemoved(Default::default()));
assert!(service_impl.update_my_infos().await);
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(SECOND_PEER_ID)
);
assert!(!service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(service_impl.route_table.peer_reachable(SECOND_PEER_ID));
}
#[tokio::test]
async fn sync_route_info_marks_credential_sender_and_filters_entries() {
let peer_mgr = create_mock_pmgr().await;
@@ -5489,7 +5935,7 @@ mod tests {
);
let mut self_info = self_info;
self_info.version = 1;
self_info.last_update = Some(SystemTime::now().into());
self_info.last_update = Some(Timestamp::now());
{
let mut guard = service_impl.synced_route_info.peer_infos.write();
guard.insert(service_impl.my_peer_id, self_info);
-3
View File
@@ -1220,9 +1220,6 @@ async fn credential_expiry_disconnects_from_all_admins() {
.await;
tokio::time::sleep(Duration::from_secs(3)).await;
admin_a
.get_global_ctx()
.issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged);
wait_for_condition(
|| {
+1
View File
@@ -1,6 +1,7 @@
use std::fmt::Display;
include!(concat!(env!("OUT_DIR"), "/acl.rs"));
include!(concat!(env!("OUT_DIR"), "/acl.serde.rs"));
impl Acl {
pub fn is_empty(&self) -> bool {
+5
View File
@@ -1,5 +1,7 @@
pub mod config {
include!(concat!(env!("OUT_DIR"), "/api.config.rs"));
include!(concat!(env!("OUT_DIR"), "/api.config.serde.rs"));
pub struct Patchable<T> {
pub action: Option<ConfigPatchAction>,
pub value: Option<T>,
@@ -77,6 +79,7 @@ pub mod config {
pub mod instance {
include!(concat!(env!("OUT_DIR"), "/api.instance.rs"));
include!(concat!(env!("OUT_DIR"), "/api.instance.serde.rs"));
impl PeerRoutePair {
pub fn get_latency_ms(&self) -> Option<f64> {
@@ -229,10 +232,12 @@ pub mod instance {
pub mod logger {
include!(concat!(env!("OUT_DIR"), "/api.logger.rs"));
include!(concat!(env!("OUT_DIR"), "/api.logger.serde.rs"));
}
pub mod manage {
include!(concat!(env!("OUT_DIR"), "/api.manage.rs"));
include!(concat!(env!("OUT_DIR"), "/api.manage.serde.rs"));
}
#[cfg(test)]
+2 -1
View File
@@ -16,7 +16,7 @@ enum NetworkingMethod {
enum ConfigSource {
ConfigSourceUnspecified = 0;
ConfigSourceUser = 1;
ConfigSourceWebhook = 2;
ConfigSourceWeb = 2;
}
message NetworkConfig {
@@ -101,6 +101,7 @@ message NetworkConfig {
optional string ipv6_public_addr_prefix = 64;
optional bool disable_relay_data = 65;
optional bool enable_udp_broadcast_relay = 66;
optional uint32 socket_mark = 67;
}
message PortForwardConfig {
+7
View File
@@ -77,6 +77,13 @@ message FlagsInConfig {
bool disable_upnp = 40;
bool disable_relay_data = 41;
bool enable_udp_broadcast_relay = 42;
// Linux-only: SO_MARK (fwmark) value applied to every outbound underlay
// socket (TCP/UDP/QUIC/WS/WG connectors and listeners). Unset = leave
// SO_MARK untouched (kernel default 0). Any set value (including 0) is
// applied via setsockopt. Requires CAP_NET_ADMIN; silently ignored on
// non-Linux platforms.
optional uint32 socket_mark = 43;
}
message RpcDescriptor {
+14 -3
View File
@@ -1,15 +1,26 @@
use anyhow::Context;
use base64::{Engine as _, prelude::BASE64_STANDARD};
use std::time::SystemTime;
use std::{
fmt::{self, Display},
str::FromStr,
};
use anyhow::Context;
use base64::{Engine as _, prelude::BASE64_STANDARD};
use strum::VariantArray;
use crate::tunnel::{IpScheme, packet_def::CompressorAlgo};
include!(concat!(env!("OUT_DIR"), "/common.rs"));
include!(concat!(env!("OUT_DIR"), "/common.serde.rs"));
pub trait TimestampExt {
fn now() -> Self;
}
impl TimestampExt for prost_wkt_types::Timestamp {
fn now() -> Self {
SystemTime::now().into()
}
}
impl From<uuid::Uuid> for Uuid {
fn from(uuid: uuid::Uuid) -> Self {
+5 -10
View File
@@ -1,10 +1,9 @@
#![allow(clippy::module_inception)]
use prost::DecodeError;
use super::rpc_types;
include!(concat!(env!("OUT_DIR"), "/error.rs"));
include!(concat!(env!("OUT_DIR"), "/error.serde.rs"));
impl From<&rpc_types::error::Error> for Error {
fn from(e: &rpc_types::error::Error) -> Self {
@@ -15,10 +14,10 @@ impl From<&rpc_types::error::Error> for Error {
error_message: format!("{:?}", e),
})),
},
rpc_types::error::Error::DecodeError(_) => Self {
rpc_types::error::Error::DecodeError => Self {
error_kind: Some(ProtoError::ProstDecodeError(ProstDecodeError {})),
},
rpc_types::error::Error::EncodeError(_) => Self {
rpc_types::error::Error::EncodeError => Self {
error_kind: Some(ProtoError::ProstEncodeError(ProstEncodeError {})),
},
rpc_types::error::Error::InvalidMethodIndex(m, s) => Self {
@@ -59,12 +58,8 @@ impl From<&Error> for rpc_types::error::Error {
Some(ProtoError::ExecuteError(e)) => {
Self::ExecutionError(anyhow::anyhow!(e.error_message.clone()))
}
Some(ProtoError::ProstDecodeError(_)) => {
Self::DecodeError(DecodeError::new("decode error"))
}
Some(ProtoError::ProstEncodeError(_)) => {
Self::DecodeError(DecodeError::new("encode error"))
}
Some(ProtoError::ProstDecodeError(_)) => Self::DecodeError,
Some(ProtoError::ProstEncodeError(_)) => Self::EncodeError,
Some(ProtoError::InvalidMethodIndex(e)) => {
Self::InvalidMethodIndex(e.method_index as u8, e.service_name.clone())
}

Some files were not shown because too many files have changed in this diff Show More