diff --git a/CONTEXT.md b/CONTEXT.md index 83652d1a..66e647a9 100644 --- a/CONTEXT.md +++ b/CONTEXT.md @@ -22,6 +22,25 @@ operation transition. Host capability operations use a separate seam. They turn Host readiness into Rust task wakeups and do not share the caller-to-core broker state machine. +## Attached peer + +An attached peer is an ordinary `PeerManagerCore` connected to another +`PeerManagerCore` through an authenticated in-process transport. Each +authenticated portal client owns one complete peer manager. The managers are +protocol peers; `attached` describes only the local transport and its trusted +ingress provenance, not a parent/child peer role. + +Each manager owns its ACL execution state, route service, RPC endpoint, secure +sessions, packet processing, and lifecycle. Portal code supplies raw packets +and peer configuration but does not build, reload, or coordinate ACL filters. + +When the network manager uses Secure Mode, an attached peer authenticates as a +credential peer. Its portal-owned, in-memory credential grant carries ACL +groups and is revoked with the attached runtime; the peer never receives the +network secret or ACL group secrets. A non-Secure-Mode network retains the +legacy admin-attached identity for compatibility. A credential peer cannot host +a portal because it cannot issue credential grants. + ## Compact compatibility Host A compact compatibility Host retains accepted values in the authoritative TOML diff --git a/Cargo.lock b/Cargo.lock index cf9b0fa8..534dae33 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2336,6 +2336,7 @@ dependencies = [ "hickory-proto", "hickory-resolver", "hickory-server", + "hkdf", "http", "humansize", "humantime-serde", @@ -2372,6 +2373,7 @@ dependencies = [ "serde_json", "serial_test", "service-manager", + "sha2", "shellexpand", "smoltcp", "socket2 0.5.10", diff --git a/README.md b/README.md index e884bffa..9754b6bc 100644 --- a/README.md +++ b/README.md @@ -252,8 +252,12 @@ ios <-.-> nodea <--> nodeb <-.-> id1 1. Start EasyTier with WireGuard portal enabled: ```bash -# Listen on 0.0.0.0:11013 and use 10.14.14.0/24 subnet for WireGuard clients -sudo easytier-core -i 10.144.144.1 --vpn-portal wg://0.0.0.0:11013/10.14.14.0/24 +# Register one WireGuard client as virtual peer 10.144.144.3 +sudo easytier-core -i 10.144.144.1 \ + --network-secret portal-secret \ + --vpn-portal wg://0.0.0.0:11013 \ + --vpn-portal-private-key "$(wg genkey)" \ + --vpn-portal-client phone=10.144.144.3 ``` 2. Get WireGuard client configuration: @@ -263,10 +267,10 @@ sudo easytier-core -i 10.144.144.1 --vpn-portal wg://0.0.0.0:11013/10.14.14.0/24 easytier-cli vpn-portal ``` -3. In the output configuration: - - Set `Interface.Address` to an available IP from the WireGuard subnet - - Set `Peer.Endpoint` to the public IP/domain of your EasyTier node - - Import the modified configuration into your WireGuard client +3. In the output configuration, replace a wildcard `Peer.Endpoint` with the + public IP/domain of your EasyTier node, then import it. `Interface.Address` + is local to that WireGuard client and may be changed to any IPv4 address; + EasyTier translates it to the registered virtual-peer address. #### Self-Hosted Public Shared Node diff --git a/README_CN.md b/README_CN.md index 3d22be93..2e2011ee 100644 --- a/README_CN.md +++ b/README_CN.md @@ -250,8 +250,12 @@ ios <-.-> nodea <--> nodeb <-.-> id1 1. 启动启用 WireGuard 门户的 EasyTier: ```bash -# 在 0.0.0.0:11013 上监听,并使用 10.14.14.0/24 子网作为 WireGuard 客户端 -sudo easytier-core -i 10.144.144.1 --vpn-portal wg://0.0.0.0:11013/10.14.14.0/24 +# 将一个 WireGuard 客户端注册为虚拟 peer 10.144.144.3 +sudo easytier-core -i 10.144.144.1 \ + --network-secret portal-secret \ + --vpn-portal wg://0.0.0.0:11013 \ + --vpn-portal-private-key "$(wg genkey)" \ + --vpn-portal-client phone=10.144.144.3 ``` 2. 获取 WireGuard 客户端配置: @@ -261,10 +265,9 @@ sudo easytier-core -i 10.144.144.1 --vpn-portal wg://0.0.0.0:11013/10.14.14.0/24 easytier-cli vpn-portal ``` -3. 在输出配置中: - - 将 `Interface.Address` 设置为 WireGuard 子网中的可用 IP - - 将 `Peer.Endpoint` 设置为您的 EasyTier 节点的公网 IP/域名 - - 将修改后的配置导入到您的 WireGuard 客户端 +3. 如果输出配置中的 `Peer.Endpoint` 是通配地址,将其替换为 EasyTier + 节点的公网 IP/域名后即可导入。`Interface.Address` 只是客户端本地地址, + 可以改为任意 IPv4 地址;EasyTier 会把它转换成已注册的虚拟 peer 地址。 #### 自建公共共享节点 diff --git a/easytier-core/src/config/api.rs b/easytier-core/src/config/api.rs index 515b3af1..c7411db0 100644 --- a/easytier-core/src/config/api.rs +++ b/easytier-core/src/config/api.rs @@ -80,11 +80,19 @@ pub fn network_config_from_toml(config: &TomlConfig) -> NetworkConfig { } if let Some(vpn_config) = config.get_vpn_portal_config() { - result.enable_vpn_portal = Some(true); - result.vpn_portal_client_network_addr = - Some(vpn_config.client_cidr.first_address().to_string()); - result.vpn_portal_client_network_len = Some(vpn_config.client_cidr.network_length() as i32); - result.vpn_portal_listen_port = Some(vpn_config.wireguard_listen.port() as i32); + result.vpn_portal_config = Some(manage::VpnPortalConfig { + wireguard_listen: vpn_config.wireguard_listen.to_string(), + wireguard_private_key: vpn_config.wireguard_private_key, + clients: vpn_config + .clients + .into_iter() + .map(|client| manage::VpnPortalClientConfig { + name: client.name, + virtual_ip: client.virtual_ip.to_string(), + groups: client.groups, + }) + .collect(), + }); } if let Some(routes) = config.get_routes() diff --git a/easytier-core/src/config/api_input.rs b/easytier-core/src/config/api_input.rs index 281985cc..489fce07 100644 --- a/easytier-core/src/config/api_input.rs +++ b/easytier-core/src/config/api_input.rs @@ -9,7 +9,7 @@ use crate::config::{ MappedListenerPolicy, normalize_secure_mode_config, toml::{ ConfigLoader, NetworkIdentity, PeerConfig, PortForwardConfig, TomlConfigLoader, - VpnPortalConfig, gen_default_flags, + VpnPortalClientConfig, VpnPortalConfig, gen_default_flags, }, }; @@ -98,6 +98,7 @@ fn parse_peer_urls(peer_urls: &[String]) -> Result, anyhow::Erro } impl NetworkConfigExt for NetworkConfig { + #[allow(deprecated)] fn gen_config(&self) -> Result { let cfg = TomlConfigLoader::default(); cfg.set_id( @@ -219,29 +220,37 @@ impl NetworkConfigExt for NetworkConfig { ); } - if self.enable_vpn_portal.unwrap_or_default() { - let cidr = format!( - "{}/{}", - self.vpn_portal_client_network_addr - .clone() - .unwrap_or_default(), - self.vpn_portal_client_network_len.unwrap_or(24) + if self.enable_vpn_portal == Some(true) { + anyhow::bail!( + "legacy VPN portal configuration is no longer supported; configure vpn_portal_config with named clients" ); + } + + if let Some(vpn_config) = &self.vpn_portal_config { cfg.set_vpn_portal_config(VpnPortalConfig { - client_cidr: cidr - .parse() - .with_context(|| format!("failed to parse vpn portal client cidr: {}", cidr))?, - wireguard_listen: format!( - "0.0.0.0:{}", - self.vpn_portal_listen_port.unwrap_or_default() - ) - .parse() - .with_context(|| { + wireguard_listen: vpn_config.wireguard_listen.parse().with_context(|| { format!( - "failed to parse vpn portal wireguard listen port. {:?}", - self.vpn_portal_listen_port + "failed to parse vpn portal wireguard listen address: {}", + vpn_config.wireguard_listen ) })?, + wireguard_private_key: vpn_config.wireguard_private_key.clone(), + clients: vpn_config + .clients + .iter() + .map(|client| { + Ok(VpnPortalClientConfig { + name: client.name.clone(), + virtual_ip: client.virtual_ip.parse().with_context(|| { + format!( + "failed to parse vpn portal virtual IP for client {}: {}", + client.name, client.virtual_ip + ) + })?, + groups: client.groups.clone(), + }) + }) + .collect::, anyhow::Error>>()?, }); } @@ -556,13 +565,19 @@ impl NetworkConfigExt for NetworkConfig { } if let Some(vpn_config) = config.get_vpn_portal_config() { - result.enable_vpn_portal = Some(true); - - let cidr = vpn_config.client_cidr; - result.vpn_portal_client_network_addr = Some(cidr.first_address().to_string()); - result.vpn_portal_client_network_len = Some(cidr.network_length() as i32); - - result.vpn_portal_listen_port = Some(vpn_config.wireguard_listen.port() as i32); + result.vpn_portal_config = Some(manage::VpnPortalConfig { + wireguard_listen: vpn_config.wireguard_listen.to_string(), + wireguard_private_key: vpn_config.wireguard_private_key, + clients: vpn_config + .clients + .into_iter() + .map(|client| manage::VpnPortalClientConfig { + name: client.name, + virtual_ip: client.virtual_ip.to_string(), + groups: client.groups, + }) + .collect(), + }); } if let Some(routes) = config.get_routes() @@ -650,3 +665,80 @@ impl NetworkConfigExt for NetworkConfig { Ok(result) } } + +#[cfg(test)] +mod tests { + #![allow(deprecated)] + + use super::*; + + fn api_portal_config() -> manage::VpnPortalConfig { + manage::VpnPortalConfig { + wireguard_listen: "0.0.0.0:51820".to_owned(), + wireguard_private_key: Some("server-private-key".to_owned()), + clients: vec![manage::VpnPortalClientConfig { + name: "alice".to_owned(), + virtual_ip: "10.144.144.10".to_owned(), + groups: vec!["staff".to_owned()], + }], + } + } + + fn standalone_config() -> NetworkConfig { + NetworkConfig { + networking_method: Some(NetworkingMethod::Standalone as i32), + ..Default::default() + } + } + + #[test] + fn vpn_portal_api_config_round_trips_through_toml_model() { + let input = NetworkConfig { + vpn_portal_config: Some(api_portal_config()), + ..standalone_config() + }; + + let config = input.gen_config().unwrap(); + let portal = config.get_vpn_portal_config().unwrap(); + assert_eq!(portal.wireguard_listen, "0.0.0.0:51820".parse().unwrap()); + assert_eq!( + portal.wireguard_private_key.as_deref(), + Some("server-private-key") + ); + assert_eq!(portal.clients[0].name, "alice"); + assert_eq!(portal.clients[0].virtual_ip.to_string(), "10.144.144.10"); + assert_eq!(portal.clients[0].groups, vec!["staff".to_owned()]); + + let output = NetworkConfig::new_from_config(&config).unwrap(); + assert_eq!(output.vpn_portal_config, input.vpn_portal_config); + assert_eq!(output.enable_vpn_portal, None); + } + + #[test] + fn legacy_enabled_vpn_portal_config_reports_migration_error() { + let error = NetworkConfig { + enable_vpn_portal: Some(true), + ..standalone_config() + } + .gen_config() + .unwrap_err() + .to_string(); + + assert!(error.contains("legacy VPN portal"), "{error}"); + } + + #[test] + fn legacy_disabled_vpn_portal_defaults_are_ignored() { + let config = NetworkConfig { + enable_vpn_portal: Some(false), + vpn_portal_listen_port: Some(0), + vpn_portal_client_network_addr: Some(String::new()), + vpn_portal_client_network_len: Some(0), + ..standalone_config() + } + .gen_config() + .unwrap(); + + assert!(config.get_vpn_portal_config().is_none()); + } +} diff --git a/easytier-core/src/config/peers.rs b/easytier-core/src/config/peers.rs index e2ff3d69..ee40002d 100644 --- a/easytier-core/src/config/peers.rs +++ b/easytier-core/src/config/peers.rs @@ -5,7 +5,7 @@ //! `crate::peers`. use anyhow::Context as _; -use cidr::{Ipv4Cidr, Ipv6Cidr}; +use cidr::Ipv6Cidr; use easytier_proto::common::{FlagsInConfig, PeerFeatureFlag, SecureModeConfig, StunInfo}; use serde::{Deserialize, Serialize}; @@ -185,6 +185,14 @@ impl AclRuleConfig { Ok(()) } + pub(crate) fn for_credential_peer(&self) -> Self { + let mut config = self.clone(); + if let Some(acl) = config.acl.as_mut().and_then(|acl| acl.acl_v1.as_mut()) { + acl.group = None; + } + config + } + pub fn build(&self) -> anyhow::Result> { let mut config = self.clone(); config.generate_acl_from_whitelists()?; @@ -229,7 +237,6 @@ pub struct PeerRuntimeSnapshot { pub easytier_version: String, pub avoid_relay_data_preference: bool, pub flags: FlagsInConfig, - pub vpn_portal_cidr: Option, pub pinned_peers: Vec<(url::Url, Option)>, pub peer_group_memberships: Vec, pub acl_group_declarations: Vec, @@ -246,7 +253,6 @@ impl PeerRuntimeSnapshot { easytier_version: env!("CARGO_PKG_VERSION").to_owned(), avoid_relay_data_preference, flags, - vpn_portal_cidr: None, pinned_peers: Vec::new(), peer_group_memberships: Vec::new(), acl_group_declarations: Vec::new(), @@ -312,4 +318,42 @@ mod tests { assert!(error.to_string().contains("Start port must be <= end port")); } + + #[test] + fn credential_peer_acl_preserves_chains_without_group_secrets() { + let config = AclRuleConfig { + acl: Some(Acl { + acl_v1: Some(AclV1 { + chains: vec![Chain { + name: "forward".to_owned(), + chain_type: ChainType::Forward as i32, + rules: vec![Rule { + action: Action::Drop as i32, + ..Default::default() + }], + ..Default::default() + }], + group: Some(GroupInfo { + declares: vec![crate::proto::acl::GroupIdentity { + group_name: "ops".to_owned(), + group_secret: "secret".to_owned(), + }], + members: vec!["ops".to_owned()], + }), + }), + }), + tcp_whitelist: vec!["22".to_owned()], + ..Default::default() + }; + + let sanitized = config.for_credential_peer(); + + let acl = sanitized.acl.unwrap().acl_v1.unwrap(); + assert_eq!(acl.chains.len(), 1); + assert_eq!(acl.chains[0].name, "forward"); + assert_eq!(acl.chains[0].rules[0].action, Action::Drop as i32); + assert!(acl.group.is_none()); + assert_eq!(sanitized.tcp_whitelist, ["22"]); + assert!(config.acl.unwrap().acl_v1.unwrap().group.is_some()); + } } diff --git a/easytier-core/src/config/runtime.rs b/easytier-core/src/config/runtime.rs index 6b2ae996..d7b5cf6e 100644 --- a/easytier-core/src/config/runtime.rs +++ b/easytier-core/src/config/runtime.rs @@ -145,6 +145,15 @@ impl CoreRuntimeConfigStore { pub fn subscribe_service_runtime_changes(&self) -> tokio::sync::watch::Receiver { self.inner.service_changes.subscribe() } + #[cfg(test)] + pub(crate) fn peer_change_subscriber_count(&self) -> usize { + self.inner.peer_changes.receiver_count() + } + + #[cfg(test)] + pub(crate) fn service_change_subscriber_count(&self) -> usize { + self.inner.service_changes.receiver_count() + } } #[cfg(test)] diff --git a/easytier-core/src/config/toml.rs b/easytier-core/src/config/toml.rs index 2b0c2b53..ef8ca432 100644 --- a/easytier-core/src/config/toml.rs +++ b/easytier-core/src/config/toml.rs @@ -276,6 +276,9 @@ pub trait ConfigLoader: Send + Sync { fn set_network_config_source(&self, _source: Option) {} fn dump(&self) -> String; + fn dump_redacted(&self) -> String { + self.dump() + } } pub trait LoggingConfigLoader { @@ -435,10 +438,37 @@ impl LoggingConfigLoader for &LoggingConfig { } } -#[derive(Debug, Clone, Deserialize, Serialize, PartialEq)] +#[derive(Clone, Deserialize, Serialize, PartialEq)] +#[serde(deny_unknown_fields)] pub struct VpnPortalConfig { - pub client_cidr: cidr::Ipv4Cidr, pub wireguard_listen: SocketAddr, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub wireguard_private_key: Option, + #[serde(default)] + pub clients: Vec, +} + +impl std::fmt::Debug for VpnPortalConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("VpnPortalConfig") + .field("wireguard_listen", &self.wireguard_listen) + .field( + "wireguard_private_key", + &self.wireguard_private_key.as_ref().map(|_| ""), + ) + .field("clients", &self.clients) + .finish() + } +} + +#[derive(Debug, Clone, Deserialize, Serialize, PartialEq, Eq)] +#[serde(deny_unknown_fields)] +pub struct VpnPortalClientConfig { + pub name: String, + pub virtual_ip: std::net::Ipv4Addr, + #[serde(default)] + pub groups: Vec, } #[derive(Debug, Clone, PartialEq, Deserialize)] @@ -549,6 +579,57 @@ impl TomlConfig { } } + #[cfg(feature = "config-write")] + fn config_for_dump(&self) -> Config { + let mut config = self.config.lock().unwrap().clone(); + Self::normalize_config_source(&mut config); + config.flags = Some(flags_diff_from_default(&self.get_flags())); + config + } + + #[cfg(feature = "config-write")] + fn redact_secrets(config: &mut Config) { + const REDACTED: &str = ""; + + if let Some(secret) = config + .network_identity + .as_mut() + .and_then(|identity| identity.network_secret.as_mut()) + && !secret.is_empty() + { + *secret = REDACTED.to_owned(); + } + if let Some(private_key) = config + .secure_mode + .as_mut() + .and_then(|secure_mode| secure_mode.local_private_key.as_mut()) + && !private_key.is_empty() + { + *private_key = REDACTED.to_owned(); + } + if let Some(private_key) = config + .vpn_portal_config + .as_mut() + .and_then(|portal| portal.wireguard_private_key.as_mut()) + && !private_key.is_empty() + { + *private_key = REDACTED.to_owned(); + } + if let Some(declarations) = config + .acl + .as_mut() + .and_then(|acl| acl.acl_v1.as_mut()) + .and_then(|acl| acl.group.as_mut()) + .map(|group| &mut group.declares) + { + for declaration in declarations { + if !declaration.group_secret.is_empty() { + declaration.group_secret = REDACTED.to_owned(); + } + } + } + } + pub fn new_from_str(config_str: &str) -> Result { Self::new_from_str_with_source("inline config", config_str) } @@ -1027,9 +1108,19 @@ impl ConfigLoader for TomlConfig { fn dump(&self) -> String { #[cfg(feature = "config-write")] { - let mut config = self.config.lock().unwrap().clone(); - Self::normalize_config_source(&mut config); - config.flags = Some(flags_diff_from_default(&self.get_flags())); + toml::to_string_pretty(&self.config_for_dump()).unwrap() + } + #[cfg(not(feature = "config-write"))] + { + panic!("this build does not include TOML configuration serialization") + } + } + + fn dump_redacted(&self) -> String { + #[cfg(feature = "config-write")] + { + let mut config = self.config_for_dump(); + Self::redact_secrets(&mut config); toml::to_string_pretty(&config).unwrap() } #[cfg(not(feature = "config-write"))] @@ -1097,6 +1188,72 @@ socket_mark = 0 assert_eq!(restored.get_flags().socket_mark, Some(0)); } + #[test] + fn legacy_vpn_portal_client_cidr_is_rejected_explicitly() { + let error = TomlConfig::new_from_str( + r#" +[vpn_portal_config] +client_cidr = "10.14.14.0/24" +wireguard_listen = "0.0.0.0:51820" +"#, + ) + .unwrap_err() + .to_string(); + + assert!(error.contains("client_cidr"), "{error}"); + } + + #[cfg(feature = "config-write")] + #[test] + fn vpn_portal_round_trip_and_redacted_dump_preserve_dump_semantics() { + let config = TomlConfig::new_from_str( + r#" +[network_identity] +network_name = "network-a" +network_secret = "network-secret" + +[secure_mode] +enabled = true +local_private_key = "noise-private-key" + +[vpn_portal_config] +wireguard_listen = "0.0.0.0:51820" +wireguard_private_key = "wireguard-private-key" + +[[vpn_portal_config.clients]] +name = "alice" +virtual_ip = "10.144.144.10" +groups = ["staff"] + +[acl.acl_v1.group] + +[[acl.acl_v1.group.declares]] +group_name = "staff" +group_secret = "group-secret" +"#, + ) + .unwrap(); + + let dumped = config.dump(); + assert!(dumped.contains("network-secret")); + assert!(dumped.contains("noise-private-key")); + assert!(dumped.contains("wireguard-private-key")); + assert!(dumped.contains("group-secret")); + assert_eq!( + TomlConfig::new_from_str(&dumped) + .unwrap() + .get_vpn_portal_config(), + config.get_vpn_portal_config() + ); + + let redacted = config.dump_redacted(); + assert!(!redacted.contains("network-secret")); + assert!(!redacted.contains("noise-private-key")); + assert!(!redacted.contains("wireguard-private-key")); + assert!(!redacted.contains("group-secret")); + assert_eq!(redacted.matches("").count(), 4); + } + #[test] fn hostname_normalization_is_portable_and_has_no_host_fallback() { let absent = TomlConfig::default(); diff --git a/easytier-core/src/gateway/proxy/cidr_monitor.rs b/easytier-core/src/gateway/proxy/cidr_monitor.rs index 3e0b6747..ef3d58d9 100644 --- a/easytier-core/src/gateway/proxy/cidr_monitor.rs +++ b/easytier-core/src/gateway/proxy/cidr_monitor.rs @@ -14,14 +14,12 @@ use crate::{ #[derive(Clone, Debug, Default, PartialEq, Eq)] pub(crate) struct ProxyCidrConfigSnapshot { pub manual_routes: Option>, - pub vpn_portal_cidr: Option, } impl From<&CoreInstanceRuntimeConfig> for ProxyCidrConfigSnapshot { fn from(config: &CoreInstanceRuntimeConfig) -> Self { Self { manual_routes: config.services.manual_routes.clone(), - vpn_portal_cidr: config.peer.vpn_portal_cidr, } } } @@ -76,15 +74,12 @@ pub struct ProxyCidrDiff { } pub(crate) fn resolve_proxy_cidrs( - mut peer_routes: BTreeSet, + peer_routes: BTreeSet, config: ProxyCidrConfigSnapshot, ) -> BTreeSet { if let Some(manual_routes) = config.manual_routes { return manual_routes; } - if let Some(vpn_portal_cidr) = config.vpn_portal_cidr { - peer_routes.insert(vpn_portal_cidr); - } peer_routes } @@ -212,39 +207,34 @@ mod tests { } #[test] - fn manual_routes_override_peer_and_vpn_routes() { + fn manual_routes_override_peer_routes() { let resolved = resolve_proxy_cidrs( cidrs(&["10.0.0.0/8"]), ProxyCidrConfigSnapshot { manual_routes: Some(cidrs(&["192.0.2.0/24"])), - vpn_portal_cidr: Some("198.51.100.0/24".parse().unwrap()), }, ); assert_eq!(resolved, cidrs(&["192.0.2.0/24"])); } #[test] - fn dynamic_routes_merge_vpn_and_report_ordered_diff() { + fn dynamic_routes_report_ordered_diff() { let current = resolve_proxy_cidrs( cidrs(&["10.0.0.0/8"]), ProxyCidrConfigSnapshot { manual_routes: None, - vpn_portal_cidr: Some("192.0.2.0/24".parse().unwrap()), }, ); let diff = diff_proxy_cidrs(&cidrs(&["10.0.0.0/8", "172.16.0.0/12"]), current); - assert_eq!(diff.current, cidrs(&["10.0.0.0/8", "192.0.2.0/24"])); - assert_eq!(diff.added, vec!["192.0.2.0/24".parse().unwrap()]); + assert_eq!(diff.current, cidrs(&["10.0.0.0/8"])); + assert!(diff.added.is_empty()); assert_eq!(diff.removed, vec!["172.16.0.0/12".parse().unwrap()]); } #[test] fn runtime_store_update_changes_the_monitor_config_snapshot() { - let initial_peer = PeerRuntimeSnapshot { - vpn_portal_cidr: Some("198.51.100.0/24".parse().unwrap()), - ..Default::default() - }; + let initial_peer = PeerRuntimeSnapshot::default(); let store = CoreRuntimeConfigStore::new( CoreRuntimeConfig { manual_routes: Some(cidrs(&["192.0.2.0/24"])), @@ -254,10 +244,7 @@ mod tests { ); let initial = store.snapshot(); - let updated_peer = PeerRuntimeSnapshot { - vpn_portal_cidr: Some("203.0.113.0/24".parse().unwrap()), - ..Default::default() - }; + let updated_peer = PeerRuntimeSnapshot::default(); store.replace(CoreInstanceRuntimeConfig { services: CoreRuntimeConfig::default(), peer: Arc::new(updated_peer), @@ -270,7 +257,7 @@ mod tests { ); assert_eq!( resolve_proxy_cidrs_from_runtime(cidrs(&["10.0.0.0/8"]), updated.as_ref()), - cidrs(&["10.0.0.0/8", "203.0.113.0/24"]) + cidrs(&["10.0.0.0/8"]) ); } } diff --git a/easytier-core/src/gateway/vpn_portal.rs b/easytier-core/src/gateway/vpn_portal.rs index 353cd2bf..236ea854 100644 --- a/easytier-core/src/gateway/vpn_portal.rs +++ b/easytier-core/src/gateway/vpn_portal.rs @@ -1,713 +1,10 @@ -use std::{ - net::{IpAddr, Ipv4Addr}, - sync::Arc, +//! Protocol-neutral portal runtime and host adapter seam. + +mod ipv4_translator; +mod runtime; + +pub use runtime::{ + DEFAULT_PORTAL_CLIENT_ADDRESS, MAX_VPN_PORTAL_CLIENTS, PortalClientConfig, + PortalClientConfigPlan, PortalClientInfoSnapshot, PortalClientState, PortalHost, + PortalInfoSnapshot, PortalListener, PortalModule, PortalRuntimeConfig, PortalSession, }; - -use async_trait::async_trait; -use cidr::{Ipv4Cidr, Ipv4Inet}; -use dashmap::DashMap; -use futures::StreamExt; -use tokio::sync::Mutex; -use tokio::task::JoinSet; -use tokio_util::sync::CancellationToken; - -use crate::{ - config::runtime::CoreRuntimeConfigStore, - events::{CoreEvent, CoreEventSink}, - packet::{PacketType, ZCPacket, ZCPacketType}, - peers::{ - PeerPacketFilter, - peer_manager::{PeerManagerCore, PipelineRegistrationGuard}, - }, - socket::SocketListener, - tunnel::{Tunnel, mpsc::MpscTunnel, mpsc::MpscTunnelSender}, -}; - -const IPV4_HEADER_LEN: usize = 20; - -pub struct VpnPortalClient { - endpoint_addr: Option, - value: V, -} - -impl VpnPortalClient { - pub fn endpoint_addr(&self) -> Option<&url::Url> { - self.endpoint_addr.as_ref() - } - - pub fn value(&self) -> &V { - &self.value - } -} - -pub struct VpnPortalClientTable { - entries: DashMap>>, -} - -impl Default for VpnPortalClientTable { - fn default() -> Self { - Self { - entries: DashMap::new(), - } - } -} - -impl VpnPortalClientTable { - pub fn new() -> Self { - Self::default() - } - - pub fn len(&self) -> usize { - self.entries.len() - } - - pub fn is_empty(&self) -> bool { - self.entries.is_empty() - } - - pub fn endpoint_addrs(&self) -> Vec> { - self.entries - .iter() - .map(|entry| entry.value().endpoint_addr.clone()) - .collect() - } - - pub fn route_peer_packet(&self, packet: &ZCPacket) -> VpnPortalPeerPacketRoute { - let Some(header) = packet.peer_manager_header() else { - return VpnPortalPeerPacketRoute::Pass; - }; - if header.packet_type != PacketType::Data as u8 { - return VpnPortalPeerPacketRoute::Pass; - } - - let payload = packet.payload(); - if payload.len() < IPV4_HEADER_LEN { - return VpnPortalPeerPacketRoute::Drop; - } - if payload[0] >> 4 != 4 { - return VpnPortalPeerPacketRoute::Pass; - } - let destination = ipv4_address(&payload[16..20]); - let Some(client) = self - .entries - .get(&destination) - .map(|entry| entry.value().clone()) - else { - return VpnPortalPeerPacketRoute::Pass; - }; - - VpnPortalPeerPacketRoute::Deliver { - destination, - client, - } - } - - fn insert(&self, address: Ipv4Addr, client: Arc>) { - self.entries.insert(address, client); - } - - fn remove_if_current(&self, address: &Ipv4Addr, client: &Arc>) -> bool { - let removed = self - .entries - .remove_if(address, |_, current| Arc::ptr_eq(current, client)) - .is_some(); - if self.entries.capacity() - self.entries.len() > 16 { - self.entries.shrink_to_fit(); - } - removed - } -} - -pub enum VpnPortalPeerPacketRoute { - Pass, - Drop, - Deliver { - destination: Ipv4Addr, - client: Arc>, - }, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct VpnPortalClientPacket { - pub source: Ipv4Addr, - pub destination: Ipv4Addr, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum VpnPortalClientRemoval { - NotRegistered, - Removed(Ipv4Addr), - EntryChangedOrMissing(Ipv4Addr), -} - -pub struct VpnPortalClientSession { - table: Arc>, - client: Arc>, - registered_ip: Option, -} - -pub type VpnPortalListener = Box>>; - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct VpnPortalClientConfigPlan { - pub client_cidr: Ipv4Cidr, - pub allowed_ips: Vec, - pub listener_url: url::Url, -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct VpnPortalInfoSnapshot { - pub vpn_type: String, - pub client_config: String, - pub connected_clients: Vec, -} - -#[async_trait] -pub trait VpnPortalHost: Send + Sync + 'static { - /// Creates already-listening protocol engines. Core owns accepting from the - /// returned listeners and all portable session lifecycle after this seam. - async fn start_listeners(&self) -> anyhow::Result>; - - fn name(&self) -> String; - - fn render_client_config(&self, plan: &VpnPortalClientConfigPlan) -> String; - - fn not_started_client_config(&self) -> String { - "ERROR: VPN Portal Not Started".to_owned() - } -} - -struct VpnPortalRuntime { - cancel: CancellationToken, - tasks: JoinSet<()>, - listener_urls: Vec, - _pipeline: PipelineRegistrationGuard, -} - -struct VpnPortalSessionEventGuard { - events: Arc, - portal: String, - client: String, -} - -impl Drop for VpnPortalSessionEventGuard { - fn drop(&mut self) { - self.events.emit(CoreEvent::VpnPortalClientDisconnected { - portal: self.portal.clone(), - client: self.client.clone(), - }); - } -} - -struct VpnPortalPeerPacketFilter { - clients: Arc>, -} - -#[async_trait] -impl PeerPacketFilter for VpnPortalPeerPacketFilter { - async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option { - let client = match self.clients.route_peer_packet(&packet) { - VpnPortalPeerPacketRoute::Pass => return Some(packet), - VpnPortalPeerPacketRoute::Drop => return None, - VpnPortalPeerPacketRoute::Deliver { client, .. } => client, - }; - - let payload_offset = packet.payload_offset(); - let packet = - ZCPacket::new_from_buf(packet.inner().split_off(payload_offset), ZCPacketType::WG); - if let Err(error) = client.value().try_send(packet) { - tracing::debug!(?error, "failed to send packet to VPN portal client"); - } - None - } -} - -pub struct VpnPortalModule { - operation: Mutex<()>, - peer_manager: Arc, - runtime_config: CoreRuntimeConfigStore, - host: Option>, - events: Arc, - clients: Arc>, - runtime: Mutex>, -} - -impl VpnPortalModule { - pub fn new( - peer_manager: Arc, - runtime_config: CoreRuntimeConfigStore, - host: Option>, - events: Arc, - ) -> Arc { - Arc::new(Self { - operation: Mutex::new(()), - peer_manager, - runtime_config, - host, - events, - clients: Arc::new(VpnPortalClientTable::new()), - runtime: Mutex::new(None), - }) - } - - pub async fn start(&self) -> anyhow::Result<()> { - let _operation = self.operation.lock().await; - if self.runtime.lock().await.is_some() { - return Ok(()); - } - if self - .runtime_config - .snapshot() - .peer - .vpn_portal_cidr - .is_none() - { - return Ok(()); - } - let Some(host) = self.host.as_ref() else { - return Ok(()); - }; - - let listeners = host.start_listeners().await?; - if listeners.is_empty() { - anyhow::bail!("VPN portal host returned no active listeners"); - } - - let cancel = CancellationToken::new(); - let mut tasks = JoinSet::new(); - let mut listener_urls = Vec::with_capacity(listeners.len()); - for listener in listeners { - let local_url = listener.local_url(); - listener_urls.push(local_url); - tasks.spawn(Self::run_listener( - listener, - self.peer_manager.clone(), - self.clients.clone(), - self.events.clone(), - cancel.clone(), - )); - } - let pipeline = self - .peer_manager - .add_managed_packet_process_pipeline(Box::new(VpnPortalPeerPacketFilter { - clients: self.clients.clone(), - })) - .await; - *self.runtime.lock().await = Some(VpnPortalRuntime { - cancel, - tasks, - listener_urls: listener_urls.clone(), - _pipeline: pipeline, - }); - for local_url in listener_urls { - self.events - .emit(CoreEvent::VpnPortalStarted(local_url.to_string())); - } - Ok(()) - } - - async fn run_listener( - mut listener: VpnPortalListener, - peer_manager: Arc, - clients: Arc>, - events: Arc, - cancel: CancellationToken, - ) { - let mut sessions = JoinSet::new(); - let mut accepting = true; - loop { - while sessions.try_join_next().is_some() {} - if !accepting && sessions.is_empty() { - break; - } - tokio::select! { - _ = cancel.cancelled() => { - sessions.shutdown().await; - return; - }, - accepted = listener.accept(), if accepting => { - match accepted { - Ok(tunnel) => { - sessions.spawn(Self::run_session( - tunnel, - peer_manager.clone(), - clients.clone(), - events.clone(), - )); - } - Err(error) => { - tracing::warn!(?error, "VPN portal listener stopped accepting"); - accepting = false; - } - } - } - _ = sessions.join_next(), if !sessions.is_empty() => {} - } - } - } - - async fn run_session( - tunnel: Box, - peer_manager: Arc, - clients: Arc>, - events: Arc, - ) { - let info = tunnel.info().unwrap_or_default(); - let portal = info.local_addr.clone().unwrap_or_default().to_string(); - let client = info.remote_addr.clone().unwrap_or_default().to_string(); - let endpoint = info.remote_addr.clone().map(Into::into); - let mut tunnel = MpscTunnel::new(tunnel, None); - let mut stream = tunnel.get_stream(); - - events.emit(CoreEvent::VpnPortalClientConnected { - portal: portal.clone(), - client: client.clone(), - }); - let _event_guard = VpnPortalSessionEventGuard { - events, - portal, - client, - }; - let mut session = VpnPortalClientSession::new(clients, endpoint, tunnel.get_sink()); - loop { - let message = match stream.next().await { - Some(Ok(message)) => message, - Some(Err(error)) => { - tracing::error!(?error, "failed to receive from VPN portal client"); - break; - } - None => break, - }; - - assert_eq!(message.packet_type(), ZCPacketType::WG); - let payload = message.inner(); - let Some(packet) = session.observe_ipv4_payload(&payload) else { - tracing::error!(?payload, "failed to parse VPN portal IPv4 packet"); - continue; - }; - let _ = peer_manager - .send_msg_by_ip( - ZCPacket::new_with_payload(&payload), - IpAddr::V4(packet.destination), - false, - ) - .await; - } - - match session.close() { - VpnPortalClientRemoval::Removed(address) => { - tracing::info!(?address, "removed VPN portal client from table") - } - VpnPortalClientRemoval::EntryChangedOrMissing(address) => tracing::info!( - ?address, - "VPN portal client endpoint changed; retaining replacement" - ), - VpnPortalClientRemoval::NotRegistered => {} - } - } - - pub async fn stop(&self) { - let _operation = self.operation.lock().await; - let Some(mut runtime) = self.runtime.lock().await.take() else { - return; - }; - runtime.cancel.cancel(); - while runtime.tasks.join_next().await.is_some() {} - self.clients.entries.clear(); - } - - pub async fn info_snapshot(&self) -> VpnPortalInfoSnapshot { - let Some(host) = self.host.as_ref() else { - return VpnPortalInfoSnapshot { - vpn_type: "null".to_owned(), - client_config: String::new(), - connected_clients: Vec::new(), - }; - }; - let runtime = self.runtime.lock().await; - let started = runtime.is_some(); - let listener_url = runtime - .as_ref() - .and_then(|runtime| runtime.listener_urls.first().cloned()); - drop(runtime); - let plan = match listener_url { - Some(listener_url) => self.client_config_plan(listener_url).await, - None => None, - }; - VpnPortalInfoSnapshot { - vpn_type: host.name(), - client_config: if started { - plan.as_ref() - .map_or_else(String::new, |plan| host.render_client_config(plan)) - } else { - host.not_started_client_config() - }, - connected_clients: self - .clients - .endpoint_addrs() - .into_iter() - .map(|endpoint| endpoint.map(|url| url.to_string()).unwrap_or_default()) - .collect(), - } - } - - async fn client_config_plan( - &self, - listener_url: url::Url, - ) -> Option { - let config = self.runtime_config.snapshot(); - let client_cidr = config.peer.vpn_portal_cidr?; - let routes = self.peer_manager.list_route_snapshots().await; - let mut allowed_ips = routes - .iter() - .flat_map(|route| route.proxy_cidrs.iter().cloned()) - .collect::>(); - let local_ipv4 = config - .peer - .runtime - .core - .routes - .ipv4 - .as_ref() - .and_then(|prefix| { - let IpAddr::V4(address) = prefix.address else { - return None; - }; - Ipv4Inet::new(address, prefix.prefix_len).ok() - }); - if let Some(ipv4) = routes - .iter() - .filter_map(|route| route.ipv4_addr.map(Into::into)) - .chain(local_ipv4) - .next() - { - allowed_ips.push(ipv4.network().to_string()); - } - allowed_ips.push(client_cidr.to_string()); - Some(VpnPortalClientConfigPlan { - client_cidr, - allowed_ips, - listener_url, - }) - } -} - -impl VpnPortalClientSession { - pub fn new( - table: Arc>, - endpoint_addr: Option, - value: V, - ) -> Self { - Self { - table, - client: Arc::new(VpnPortalClient { - endpoint_addr, - value, - }), - registered_ip: None, - } - } - - pub fn observe_ipv4_payload(&mut self, payload: &[u8]) -> Option { - if payload.len() < IPV4_HEADER_LEN { - return None; - } - let packet = VpnPortalClientPacket { - source: ipv4_address(&payload[12..16]), - destination: ipv4_address(&payload[16..20]), - }; - - if self.registered_ip.is_none() { - self.table.insert(packet.source, self.client.clone()); - self.registered_ip = Some(packet.source); - } - Some(packet) - } - - pub fn registered_ip(&self) -> Option { - self.registered_ip - } - - pub fn close(&mut self) -> VpnPortalClientRemoval { - let Some(address) = self.registered_ip.take() else { - return VpnPortalClientRemoval::NotRegistered; - }; - if self.table.remove_if_current(&address, &self.client) { - VpnPortalClientRemoval::Removed(address) - } else { - VpnPortalClientRemoval::EntryChangedOrMissing(address) - } - } -} - -impl Drop for VpnPortalClientSession { - fn drop(&mut self) { - let _ = self.close(); - } -} - -fn ipv4_address(bytes: &[u8]) -> Ipv4Addr { - Ipv4Addr::new(bytes[0], bytes[1], bytes[2], bytes[3]) -} - -#[cfg(test)] -mod tests { - use super::*; - - fn ipv4_payload(source: [u8; 4], destination: [u8; 4], version: u8) -> Vec { - let mut payload = vec![0u8; IPV4_HEADER_LEN]; - payload[0] = version << 4 | 5; - payload[12..16].copy_from_slice(&source); - payload[16..20].copy_from_slice(&destination); - payload - } - - fn peer_packet(payload: &[u8], packet_type: PacketType) -> ZCPacket { - let mut packet = ZCPacket::new_with_payload(payload); - packet.fill_peer_manager_hdr(1, 2, packet_type as u8); - packet - } - - #[test] - fn session_registers_first_source_and_routes_peer_packet() { - let table = Arc::new(VpnPortalClientTable::new()); - let endpoint = Some("wg://198.51.100.2:51820".parse().unwrap()); - let mut session = VpnPortalClientSession::new(table.clone(), endpoint, "client"); - - let observed = session - .observe_ipv4_payload(&ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4)) - .unwrap(); - assert_eq!(observed.source, Ipv4Addr::new(10, 10, 0, 2)); - assert_eq!(session.registered_ip(), Some(observed.source)); - - let packet = peer_packet( - &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), - PacketType::Data, - ); - let VpnPortalPeerPacketRoute::Deliver { - destination, - client, - } = table.route_peer_packet(&packet) - else { - panic!("registered destination must be delivered"); - }; - assert_eq!(destination, observed.source); - assert_eq!(client.value(), &"client"); - } - - #[test] - fn closing_old_endpoint_does_not_remove_replacement() { - let table = Arc::new(VpnPortalClientTable::new()); - let payload = ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4); - let mut old = VpnPortalClientSession::new( - table.clone(), - Some("wg://198.51.100.2:51820".parse().unwrap()), - "old", - ); - let mut replacement = VpnPortalClientSession::new( - table.clone(), - Some("wg://198.51.100.3:51820".parse().unwrap()), - "replacement", - ); - old.observe_ipv4_payload(&payload).unwrap(); - replacement.observe_ipv4_payload(&payload).unwrap(); - - assert_eq!( - old.close(), - VpnPortalClientRemoval::EntryChangedOrMissing(Ipv4Addr::new(10, 10, 0, 2)) - ); - assert_eq!(table.len(), 1); - - let routed = peer_packet( - &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), - PacketType::Data, - ); - let VpnPortalPeerPacketRoute::Deliver { client, .. } = table.route_peer_packet(&routed) - else { - panic!("replacement must remain registered"); - }; - assert_eq!(client.value(), &"replacement"); - } - - #[test] - fn dropping_old_session_does_not_remove_same_endpoint_replacement() { - let table = Arc::new(VpnPortalClientTable::new()); - let endpoint = Some("wg://198.51.100.2:51820".parse().unwrap()); - let payload = ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4); - let mut old = VpnPortalClientSession::new(table.clone(), endpoint.clone(), "old"); - let mut replacement = VpnPortalClientSession::new(table.clone(), endpoint, "replacement"); - old.observe_ipv4_payload(&payload).unwrap(); - replacement.observe_ipv4_payload(&payload).unwrap(); - - drop(old); - - let routed = peer_packet( - &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), - PacketType::Data, - ); - let VpnPortalPeerPacketRoute::Deliver { client, .. } = table.route_peer_packet(&routed) - else { - panic!("same-endpoint replacement must remain registered"); - }; - assert_eq!(client.value(), &"replacement"); - } - - #[test] - fn close_removes_matching_entry_and_non_data_packets_pass() { - let table = Arc::new(VpnPortalClientTable::new()); - let mut session = VpnPortalClientSession::new(table.clone(), None, ()); - session - .observe_ipv4_payload(&ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4)) - .unwrap(); - - let non_data = peer_packet( - &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 4), - PacketType::Ping, - ); - assert!(matches!( - table.route_peer_packet(&non_data), - VpnPortalPeerPacketRoute::Pass - )); - assert_eq!( - session.close(), - VpnPortalClientRemoval::Removed(Ipv4Addr::new(10, 10, 0, 2)) - ); - assert!(table.is_empty()); - } - - #[test] - fn dropping_session_removes_matching_entry() { - let table = Arc::new(VpnPortalClientTable::new()); - { - let mut session = VpnPortalClientSession::new(table.clone(), None, ()); - session - .observe_ipv4_payload(&ipv4_payload([10, 10, 0, 2], [10, 10, 0, 3], 4)) - .unwrap(); - assert_eq!(table.len(), 1); - } - assert!(table.is_empty()); - } - - #[test] - fn peer_route_rejects_non_ipv4_payload() { - let table = VpnPortalClientTable::<()>::new(); - let packet = peer_packet( - &ipv4_payload([10, 10, 0, 3], [10, 10, 0, 2], 6), - PacketType::Data, - ); - assert!(matches!( - table.route_peer_packet(&packet), - VpnPortalPeerPacketRoute::Pass - )); - } - - #[test] - fn peer_route_drops_short_data_payload() { - let table = VpnPortalClientTable::<()>::new(); - let packet = peer_packet(&[0u8; IPV4_HEADER_LEN - 1], PacketType::Data); - assert!(matches!( - table.route_peer_packet(&packet), - VpnPortalPeerPacketRoute::Drop - )); - } -} diff --git a/easytier-core/src/gateway/vpn_portal/ipv4_translator.rs b/easytier-core/src/gateway/vpn_portal/ipv4_translator.rs new file mode 100644 index 00000000..ed9c0469 --- /dev/null +++ b/easytier-core/src/gateway/vpn_portal/ipv4_translator.rs @@ -0,0 +1,974 @@ +use std::net::Ipv4Addr; + +const IPV4_MIN_HEADER_LEN: usize = 20; +const TCP_MIN_HEADER_LEN: usize = 20; +const UDP_HEADER_LEN: usize = 8; +const ICMP_MIN_HEADER_LEN: usize = 8; + +const IP_PROTOCOL_ICMP: u8 = 1; +const IP_PROTOCOL_TCP: u8 = 6; +const IP_PROTOCOL_UDP: u8 = 17; + +const IPV4_CHECKSUM_OFFSET: usize = 10; +const IPV4_SOURCE_OFFSET: usize = 12; +const IPV4_DESTINATION_OFFSET: usize = 16; +const TCP_CHECKSUM_OFFSET: usize = 16; +const UDP_CHECKSUM_OFFSET: usize = 6; +const ICMP_CHECKSUM_OFFSET: usize = 2; +const ICMP_QUOTED_PACKET_OFFSET: usize = 8; + +#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] +pub(crate) enum Ipv4TranslationError { + #[error("IPv4 packet is too short: expected at least 20 bytes, got {actual}")] + PacketTooShort { actual: usize }, + #[error("unsupported IP version {version}; expected IPv4")] + UnsupportedIpVersion { version: u8 }, + #[error("invalid IPv4 IHL {ihl_words}; expected at least 5 words")] + InvalidHeaderLength { ihl_words: u8 }, + #[error("truncated IPv4 header: header is {header_len} bytes, packet is {actual} bytes")] + TruncatedHeader { header_len: usize, actual: usize }, + #[error("IPv4 total length {declared} does not match payload length {actual}")] + TotalLengthMismatch { declared: usize, actual: usize }, + #[error("unexpected IPv4 source: expected {expected}, got {actual}")] + UnexpectedSource { + expected: Ipv4Addr, + actual: Ipv4Addr, + }, + #[error("unexpected IPv4 destination: expected {expected}, got {actual}")] + UnexpectedDestination { + expected: Ipv4Addr, + actual: Ipv4Addr, + }, + #[error("unsupported IPv4 protocol {protocol}")] + UnsupportedProtocol { protocol: u8 }, + #[error("truncated {protocol} header: expected at least {required} bytes, got {actual}")] + TruncatedTransportHeader { + protocol: &'static str, + required: usize, + actual: usize, + }, + #[error("invalid TCP data offset {data_offset_words}; expected at least 5 words")] + InvalidTcpHeaderLength { data_offset_words: u8 }, + #[error("truncated TCP header: header is {header_len} bytes, fragment carries {actual} bytes")] + TruncatedTcpHeader { header_len: usize, actual: usize }, + #[error("invalid UDP length {declared} for an IPv4 payload carrying {actual} UDP bytes")] + InvalidUdpLength { declared: usize, actual: usize }, + #[error("non-final IPv4 fragment carries {actual} bytes; expected a multiple of 8")] + InvalidFragmentLength { actual: usize }, + #[error("ICMP error quotes only {actual} IPv4 bytes; expected at least 20")] + QuotedPacketTooShort { actual: usize }, + #[error("ICMP error quotes IP version {version}; expected IPv4")] + UnsupportedQuotedIpVersion { version: u8 }, + #[error("ICMP error quotes an invalid IPv4 IHL {ihl_words}; expected at least 5 words")] + InvalidQuotedHeaderLength { ihl_words: u8 }, + #[error( + "ICMP error quotes a truncated IPv4 header: header is {header_len} bytes, quote is {actual} bytes" + )] + TruncatedQuotedHeader { header_len: usize, actual: usize }, + #[error( + "ICMP error quotes an IPv4 total length {declared} smaller than its {header_len}-byte header" + )] + InvalidQuotedTotalLength { declared: usize, header_len: usize }, + #[error("unsupported protocol {protocol} in translated ICMP IPv4 quote")] + UnsupportedQuotedProtocol { protocol: u8 }, +} + +#[derive(Clone, Copy)] +enum AddressField { + Source, + Destination, +} + +impl AddressField { + fn offset(self) -> usize { + match self { + Self::Source => IPV4_SOURCE_OFFSET, + Self::Destination => IPV4_DESTINATION_OFFSET, + } + } +} + +#[derive(Clone, Copy)] +struct Ipv4Layout { + header_len: usize, + protocol: u8, + fragment_offset: u16, + more_fragments: bool, +} + +#[derive(Clone, Copy)] +struct ChecksumField { + offset: usize, + udp: bool, +} + +#[derive(Clone, Copy)] +struct QuotedIpv4Plan { + header_offset: usize, + header_len: usize, + replace_source: bool, + replace_destination: bool, + transport_checksum: Option, +} + +#[derive(Clone, Copy)] +enum TransportPlan { + HeaderOnly, + WithPseudoHeaderChecksum(ChecksumField), + Icmp { + checksum_offset: usize, + quoted: Option, + }, +} + +pub(crate) fn rewrite_ipv4_source( + packet: &mut [u8], + old_address: Ipv4Addr, + new_address: Ipv4Addr, +) -> Result<(), Ipv4TranslationError> { + rewrite_ipv4_address(packet, AddressField::Source, old_address, new_address) +} + +pub(crate) fn rewrite_ipv4_destination( + packet: &mut [u8], + old_address: Ipv4Addr, + new_address: Ipv4Addr, +) -> Result<(), Ipv4TranslationError> { + rewrite_ipv4_address(packet, AddressField::Destination, old_address, new_address) +} + +fn rewrite_ipv4_address( + packet: &mut [u8], + field: AddressField, + old_address: Ipv4Addr, + new_address: Ipv4Addr, +) -> Result<(), Ipv4TranslationError> { + let layout = parse_complete_ipv4(packet)?; + let actual_address = read_ipv4_address(packet, field.offset()); + if actual_address != old_address { + return Err(match field { + AddressField::Source => Ipv4TranslationError::UnexpectedSource { + expected: old_address, + actual: actual_address, + }, + AddressField::Destination => Ipv4TranslationError::UnexpectedDestination { + expected: old_address, + actual: actual_address, + }, + }); + } + + let transport_plan = analyze_transport(packet, layout, old_address)?; + + match transport_plan { + TransportPlan::HeaderOnly => {} + TransportPlan::WithPseudoHeaderChecksum(checksum) => { + rewrite_pseudo_header_checksum(packet, checksum, old_address, new_address); + } + TransportPlan::Icmp { + checksum_offset, + quoted, + } => { + let mut fragmented_checksum = read_u16(packet, checksum_offset); + if let Some(quoted) = quoted { + rewrite_quoted_ipv4( + packet, + quoted, + old_address, + new_address, + &mut fragmented_checksum, + ); + } + + if layout.more_fragments { + write_u16(packet, checksum_offset, fragmented_checksum); + } else { + let icmp = &packet[layout.header_len..]; + let checksum = checksum_with_zeroed_word(icmp, ICMP_CHECKSUM_OFFSET); + write_u16(packet, checksum_offset, checksum); + } + } + } + + packet[field.offset()..field.offset() + 4].copy_from_slice(&new_address.octets()); + let checksum = checksum_with_zeroed_word(&packet[..layout.header_len], IPV4_CHECKSUM_OFFSET); + write_u16(packet, IPV4_CHECKSUM_OFFSET, checksum); + Ok(()) +} + +fn parse_complete_ipv4(packet: &[u8]) -> Result { + if packet.len() < IPV4_MIN_HEADER_LEN { + return Err(Ipv4TranslationError::PacketTooShort { + actual: packet.len(), + }); + } + + let version = packet[0] >> 4; + if version != 4 { + return Err(Ipv4TranslationError::UnsupportedIpVersion { version }); + } + let ihl_words = packet[0] & 0x0f; + if ihl_words < 5 { + return Err(Ipv4TranslationError::InvalidHeaderLength { ihl_words }); + } + let header_len = usize::from(ihl_words) * 4; + if header_len > packet.len() { + return Err(Ipv4TranslationError::TruncatedHeader { + header_len, + actual: packet.len(), + }); + } + + let declared = usize::from(read_u16(packet, 2)); + if declared != packet.len() { + return Err(Ipv4TranslationError::TotalLengthMismatch { + declared, + actual: packet.len(), + }); + } + + let fragment = read_u16(packet, 6); + let layout = Ipv4Layout { + header_len, + protocol: packet[9], + fragment_offset: fragment & 0x1fff, + more_fragments: fragment & 0x2000 != 0, + }; + let fragment_payload_len = packet.len() - header_len; + if layout.more_fragments && !fragment_payload_len.is_multiple_of(8) { + return Err(Ipv4TranslationError::InvalidFragmentLength { + actual: fragment_payload_len, + }); + } + Ok(layout) +} + +fn analyze_transport( + packet: &[u8], + layout: Ipv4Layout, + old_address: Ipv4Addr, +) -> Result { + if !matches!( + layout.protocol, + IP_PROTOCOL_TCP | IP_PROTOCOL_UDP | IP_PROTOCOL_ICMP + ) { + return Err(Ipv4TranslationError::UnsupportedProtocol { + protocol: layout.protocol, + }); + } + if layout.fragment_offset != 0 { + return Ok(TransportPlan::HeaderOnly); + } + + let transport_len = packet.len() - layout.header_len; + match layout.protocol { + IP_PROTOCOL_TCP => { + let required = if layout.more_fragments { + TCP_CHECKSUM_OFFSET + 2 + } else { + TCP_MIN_HEADER_LEN + }; + if transport_len < required { + return Err(Ipv4TranslationError::TruncatedTransportHeader { + protocol: "TCP", + required, + actual: transport_len, + }); + } + let data_offset_words = packet[layout.header_len + 12] >> 4; + if data_offset_words < 5 { + return Err(Ipv4TranslationError::InvalidTcpHeaderLength { data_offset_words }); + } + let tcp_header_len = usize::from(data_offset_words) * 4; + if !layout.more_fragments && tcp_header_len > transport_len { + return Err(Ipv4TranslationError::TruncatedTcpHeader { + header_len: tcp_header_len, + actual: transport_len, + }); + } + Ok(TransportPlan::WithPseudoHeaderChecksum(ChecksumField { + offset: layout.header_len + TCP_CHECKSUM_OFFSET, + udp: false, + })) + } + IP_PROTOCOL_UDP => { + if transport_len < UDP_HEADER_LEN { + return Err(Ipv4TranslationError::TruncatedTransportHeader { + protocol: "UDP", + required: UDP_HEADER_LEN, + actual: transport_len, + }); + } + let udp_len = usize::from(read_u16(packet, layout.header_len + 4)); + let invalid = udp_len < UDP_HEADER_LEN + || (!layout.more_fragments && udp_len != transport_len) + || (layout.more_fragments && udp_len <= transport_len); + if invalid { + return Err(Ipv4TranslationError::InvalidUdpLength { + declared: udp_len, + actual: transport_len, + }); + } + Ok(TransportPlan::WithPseudoHeaderChecksum(ChecksumField { + offset: layout.header_len + UDP_CHECKSUM_OFFSET, + udp: true, + })) + } + IP_PROTOCOL_ICMP => { + if transport_len < ICMP_MIN_HEADER_LEN { + return Err(Ipv4TranslationError::TruncatedTransportHeader { + protocol: "ICMP", + required: ICMP_MIN_HEADER_LEN, + actual: transport_len, + }); + } + let icmp_offset = layout.header_len; + let quoted = if is_icmp_error(packet[icmp_offset]) { + Some(analyze_quoted_ipv4( + packet, + icmp_offset + ICMP_QUOTED_PACKET_OFFSET, + old_address, + )?) + } else { + None + }; + Ok(TransportPlan::Icmp { + checksum_offset: icmp_offset + ICMP_CHECKSUM_OFFSET, + quoted, + }) + } + _ => unreachable!("supported protocol checked above"), + } +} + +fn analyze_quoted_ipv4( + packet: &[u8], + header_offset: usize, + old_address: Ipv4Addr, +) -> Result { + let quote = &packet[header_offset..]; + if quote.len() < IPV4_MIN_HEADER_LEN { + return Err(Ipv4TranslationError::QuotedPacketTooShort { + actual: quote.len(), + }); + } + let version = quote[0] >> 4; + if version != 4 { + return Err(Ipv4TranslationError::UnsupportedQuotedIpVersion { version }); + } + let ihl_words = quote[0] & 0x0f; + if ihl_words < 5 { + return Err(Ipv4TranslationError::InvalidQuotedHeaderLength { ihl_words }); + } + let header_len = usize::from(ihl_words) * 4; + if header_len > quote.len() { + return Err(Ipv4TranslationError::TruncatedQuotedHeader { + header_len, + actual: quote.len(), + }); + } + let total_len = usize::from(read_u16(quote, 2)); + if total_len < header_len { + return Err(Ipv4TranslationError::InvalidQuotedTotalLength { + declared: total_len, + header_len, + }); + } + + let replace_source = read_ipv4_address(quote, IPV4_SOURCE_OFFSET) == old_address; + let replace_destination = read_ipv4_address(quote, IPV4_DESTINATION_OFFSET) == old_address; + let fragment_offset = read_u16(quote, 6) & 0x1fff; + let visible_len = total_len.min(quote.len()); + let protocol = quote[9]; + let transport_checksum = if (!replace_source && !replace_destination) || fragment_offset != 0 { + None + } else { + let relative_checksum_offset = match protocol { + IP_PROTOCOL_TCP => header_len + TCP_CHECKSUM_OFFSET, + IP_PROTOCOL_UDP => header_len + UDP_CHECKSUM_OFFSET, + IP_PROTOCOL_ICMP => usize::MAX, + _ => { + return Err(Ipv4TranslationError::UnsupportedQuotedProtocol { protocol }); + } + }; + let checksum_visible = + relative_checksum_offset != usize::MAX && relative_checksum_offset + 2 <= visible_len; + checksum_visible.then_some(ChecksumField { + offset: header_offset + relative_checksum_offset, + udp: protocol == IP_PROTOCOL_UDP, + }) + }; + + Ok(QuotedIpv4Plan { + header_offset, + header_len, + replace_source, + replace_destination, + transport_checksum, + }) +} + +fn rewrite_quoted_ipv4( + packet: &mut [u8], + plan: QuotedIpv4Plan, + old_address: Ipv4Addr, + new_address: Ipv4Addr, + outer_icmp_checksum: &mut u16, +) { + if !plan.replace_source && !plan.replace_destination { + return; + } + + if let Some(checksum) = plan.transport_checksum { + let current = read_u16(packet, checksum.offset); + if !checksum.udp || current != 0 { + let mut updated = current; + if plan.replace_source { + updated = update_checksum_for_address(updated, old_address, new_address); + } + if plan.replace_destination { + updated = update_checksum_for_address(updated, old_address, new_address); + } + if checksum.udp && updated == 0 { + updated = u16::MAX; + } + write_tracked_word(packet, checksum.offset, updated, outer_icmp_checksum); + } + } + + if plan.replace_source { + write_tracked_address( + packet, + plan.header_offset + IPV4_SOURCE_OFFSET, + new_address, + outer_icmp_checksum, + ); + } + if plan.replace_destination { + write_tracked_address( + packet, + plan.header_offset + IPV4_DESTINATION_OFFSET, + new_address, + outer_icmp_checksum, + ); + } + + let inner_header = &packet[plan.header_offset..plan.header_offset + plan.header_len]; + let checksum = checksum_with_zeroed_word(inner_header, IPV4_CHECKSUM_OFFSET); + write_tracked_word( + packet, + plan.header_offset + IPV4_CHECKSUM_OFFSET, + checksum, + outer_icmp_checksum, + ); +} + +fn rewrite_pseudo_header_checksum( + packet: &mut [u8], + checksum: ChecksumField, + old_address: Ipv4Addr, + new_address: Ipv4Addr, +) { + let current = read_u16(packet, checksum.offset); + if checksum.udp && current == 0 { + return; + } + let mut updated = update_checksum_for_address(current, old_address, new_address); + if checksum.udp && updated == 0 { + updated = u16::MAX; + } + write_u16(packet, checksum.offset, updated); +} + +fn write_tracked_address( + packet: &mut [u8], + offset: usize, + address: Ipv4Addr, + enclosing_checksum: &mut u16, +) { + let octets = address.octets(); + write_tracked_word( + packet, + offset, + u16::from_be_bytes([octets[0], octets[1]]), + enclosing_checksum, + ); + write_tracked_word( + packet, + offset + 2, + u16::from_be_bytes([octets[2], octets[3]]), + enclosing_checksum, + ); +} + +fn write_tracked_word( + packet: &mut [u8], + offset: usize, + new_value: u16, + enclosing_checksum: &mut u16, +) { + let old_value = read_u16(packet, offset); + if old_value == new_value { + return; + } + *enclosing_checksum = update_checksum_word(*enclosing_checksum, old_value, new_value); + write_u16(packet, offset, new_value); +} + +fn update_checksum_for_address(checksum: u16, old_address: Ipv4Addr, new_address: Ipv4Addr) -> u16 { + let old = old_address.octets(); + let new = new_address.octets(); + let checksum = update_checksum_word( + checksum, + u16::from_be_bytes([old[0], old[1]]), + u16::from_be_bytes([new[0], new[1]]), + ); + update_checksum_word( + checksum, + u16::from_be_bytes([old[2], old[3]]), + u16::from_be_bytes([new[2], new[3]]), + ) +} + +fn update_checksum_word(checksum: u16, old_value: u16, new_value: u16) -> u16 { + let mut sum = u32::from(!checksum) + u32::from(!old_value) + u32::from(new_value); + while sum >> 16 != 0 { + sum = (sum & 0xffff) + (sum >> 16); + } + !(sum as u16) +} + +fn checksum_with_zeroed_word(bytes: &[u8], zero_offset: usize) -> u16 { + let mut sum = 0u32; + for (offset, chunk) in bytes.chunks(2).enumerate() { + let byte_offset = offset * 2; + let word = if byte_offset == zero_offset { + 0 + } else if let [high, low] = chunk { + u16::from_be_bytes([*high, *low]) + } else { + u16::from(chunk[0]) << 8 + }; + sum += u32::from(word); + } + while sum >> 16 != 0 { + sum = (sum & 0xffff) + (sum >> 16); + } + !(sum as u16) +} + +fn is_icmp_error(icmp_type: u8) -> bool { + matches!(icmp_type, 3 | 4 | 5 | 11 | 12) +} + +fn read_ipv4_address(bytes: &[u8], offset: usize) -> Ipv4Addr { + Ipv4Addr::new( + bytes[offset], + bytes[offset + 1], + bytes[offset + 2], + bytes[offset + 3], + ) +} + +fn read_u16(bytes: &[u8], offset: usize) -> u16 { + u16::from_be_bytes([bytes[offset], bytes[offset + 1]]) +} + +fn write_u16(bytes: &mut [u8], offset: usize, value: u16) { + bytes[offset..offset + 2].copy_from_slice(&value.to_be_bytes()); +} + +#[cfg(test)] +mod tests { + use super::*; + + const CLIENT_IP: Ipv4Addr = Ipv4Addr::new(192, 0, 2, 1); + const VIRTUAL_IP: Ipv4Addr = Ipv4Addr::new(10, 144, 144, 10); + const REMOTE_IP: Ipv4Addr = Ipv4Addr::new(10, 144, 144, 20); + + fn build_ipv4( + source: Ipv4Addr, + destination: Ipv4Addr, + protocol: u8, + payload: &[u8], + options: &[u8], + fragment: u16, + ) -> Vec { + assert_eq!(options.len() % 4, 0); + let header_len = IPV4_MIN_HEADER_LEN + options.len(); + let mut packet = vec![0; header_len + payload.len()]; + packet[0] = 0x40 | u8::try_from(header_len / 4).unwrap(); + let packet_len = u16::try_from(packet.len()).unwrap(); + write_u16(&mut packet, 2, packet_len); + write_u16(&mut packet, 4, 0x1234); + write_u16(&mut packet, 6, fragment); + packet[8] = 64; + packet[9] = protocol; + packet[IPV4_SOURCE_OFFSET..IPV4_SOURCE_OFFSET + 4].copy_from_slice(&source.octets()); + packet[IPV4_DESTINATION_OFFSET..IPV4_DESTINATION_OFFSET + 4] + .copy_from_slice(&destination.octets()); + packet[IPV4_MIN_HEADER_LEN..header_len].copy_from_slice(options); + packet[header_len..].copy_from_slice(payload); + let checksum = checksum_with_zeroed_word(&packet[..header_len], IPV4_CHECKSUM_OFFSET); + write_u16(&mut packet, IPV4_CHECKSUM_OFFSET, checksum); + packet + } + + fn tcp_segment(source: Ipv4Addr, destination: Ipv4Addr, data: &[u8]) -> Vec { + let mut tcp = vec![0; TCP_MIN_HEADER_LEN + data.len()]; + write_u16(&mut tcp, 0, 12345); + write_u16(&mut tcp, 2, 443); + tcp[12] = 5 << 4; + tcp[13] = 0x18; + write_u16(&mut tcp, 14, 4096); + tcp[TCP_MIN_HEADER_LEN..].copy_from_slice(data); + let checksum = transport_checksum(source, destination, IP_PROTOCOL_TCP, &tcp); + write_u16(&mut tcp, TCP_CHECKSUM_OFFSET, checksum); + tcp + } + + fn udp_datagram( + source: Ipv4Addr, + destination: Ipv4Addr, + data: &[u8], + checksum_enabled: bool, + ) -> Vec { + let mut udp = vec![0; UDP_HEADER_LEN + data.len()]; + write_u16(&mut udp, 0, 5353); + write_u16(&mut udp, 2, 53); + let udp_len = u16::try_from(udp.len()).unwrap(); + write_u16(&mut udp, 4, udp_len); + udp[UDP_HEADER_LEN..].copy_from_slice(data); + if checksum_enabled { + let checksum = transport_checksum(source, destination, IP_PROTOCOL_UDP, &udp); + write_u16( + &mut udp, + UDP_CHECKSUM_OFFSET, + if checksum == 0 { u16::MAX } else { checksum }, + ); + } + udp + } + + fn icmp_message(icmp_type: u8, body: &[u8]) -> Vec { + let mut icmp = vec![0; ICMP_MIN_HEADER_LEN + body.len()]; + icmp[0] = icmp_type; + icmp[1] = 0; + icmp[4..8].copy_from_slice(&[0x12, 0x34, 0, 1]); + icmp[ICMP_MIN_HEADER_LEN..].copy_from_slice(body); + let checksum = checksum_with_zeroed_word(&icmp, ICMP_CHECKSUM_OFFSET); + write_u16(&mut icmp, ICMP_CHECKSUM_OFFSET, checksum); + icmp + } + + fn transport_checksum( + source: Ipv4Addr, + destination: Ipv4Addr, + protocol: u8, + transport: &[u8], + ) -> u16 { + let mut bytes = Vec::with_capacity(12 + transport.len()); + bytes.extend_from_slice(&source.octets()); + bytes.extend_from_slice(&destination.octets()); + bytes.push(0); + bytes.push(protocol); + bytes.extend_from_slice(&u16::try_from(transport.len()).unwrap().to_be_bytes()); + bytes.extend_from_slice(transport); + checksum_with_zeroed_word( + &bytes, + 12 + if protocol == IP_PROTOCOL_TCP { + TCP_CHECKSUM_OFFSET + } else { + UDP_CHECKSUM_OFFSET + }, + ) + } + + fn assert_valid_ipv4_checksum(packet: &[u8]) { + let header_len = usize::from(packet[0] & 0x0f) * 4; + assert_eq!( + read_u16(packet, IPV4_CHECKSUM_OFFSET), + checksum_with_zeroed_word(&packet[..header_len], IPV4_CHECKSUM_OFFSET) + ); + } + + fn assert_valid_transport_checksum(packet: &[u8], protocol: u8) { + let header_len = usize::from(packet[0] & 0x0f) * 4; + let source = read_ipv4_address(packet, IPV4_SOURCE_OFFSET); + let destination = read_ipv4_address(packet, IPV4_DESTINATION_OFFSET); + let transport = &packet[header_len..]; + let offset = if protocol == IP_PROTOCOL_TCP { + TCP_CHECKSUM_OFFSET + } else { + UDP_CHECKSUM_OFFSET + }; + assert_eq!( + read_u16(transport, offset), + transport_checksum(source, destination, protocol, transport) + ); + } + + #[test] + fn rewrites_tcp_source_with_ipv4_options() { + let tcp = tcp_segment(CLIENT_IP, REMOTE_IP, b"tcp payload"); + let mut packet = build_ipv4( + CLIENT_IP, + REMOTE_IP, + IP_PROTOCOL_TCP, + &tcp, + &[1, 1, 1, 0], + 0, + ); + + rewrite_ipv4_source(&mut packet, CLIENT_IP, VIRTUAL_IP).unwrap(); + + assert_eq!(read_ipv4_address(&packet, IPV4_SOURCE_OFFSET), VIRTUAL_IP); + assert_eq!( + read_ipv4_address(&packet, IPV4_DESTINATION_OFFSET), + REMOTE_IP + ); + assert_valid_ipv4_checksum(&packet); + assert_valid_transport_checksum(&packet, IP_PROTOCOL_TCP); + } + + #[test] + fn rewrites_tcp_destination() { + let tcp = tcp_segment(REMOTE_IP, VIRTUAL_IP, b"reply"); + let mut packet = build_ipv4(REMOTE_IP, VIRTUAL_IP, IP_PROTOCOL_TCP, &tcp, &[], 0); + + rewrite_ipv4_destination(&mut packet, VIRTUAL_IP, CLIENT_IP).unwrap(); + + assert_eq!(read_ipv4_address(&packet, IPV4_SOURCE_OFFSET), REMOTE_IP); + assert_eq!( + read_ipv4_address(&packet, IPV4_DESTINATION_OFFSET), + CLIENT_IP + ); + assert_valid_ipv4_checksum(&packet); + assert_valid_transport_checksum(&packet, IP_PROTOCOL_TCP); + } + + #[test] + fn rewrites_udp_checksum_and_preserves_disabled_checksum() { + for checksum_enabled in [true, false] { + let udp = udp_datagram(CLIENT_IP, REMOTE_IP, b"dns", checksum_enabled); + let mut packet = build_ipv4(CLIENT_IP, REMOTE_IP, IP_PROTOCOL_UDP, &udp, &[], 0); + + rewrite_ipv4_source(&mut packet, CLIENT_IP, VIRTUAL_IP).unwrap(); + + assert_valid_ipv4_checksum(&packet); + let udp_offset = IPV4_MIN_HEADER_LEN + UDP_CHECKSUM_OFFSET; + if checksum_enabled { + assert_valid_transport_checksum(&packet, IP_PROTOCOL_UDP); + assert_ne!(read_u16(&packet, udp_offset), 0); + } else { + assert_eq!(read_u16(&packet, udp_offset), 0); + } + } + } + + #[test] + fn rewrites_icmp_echo_outer_address_and_checksum() { + let icmp = icmp_message(8, b"echo payload"); + let original_icmp_checksum = read_u16(&icmp, ICMP_CHECKSUM_OFFSET); + let mut packet = build_ipv4(CLIENT_IP, REMOTE_IP, IP_PROTOCOL_ICMP, &icmp, &[], 0); + + rewrite_ipv4_source(&mut packet, CLIENT_IP, VIRTUAL_IP).unwrap(); + + assert_valid_ipv4_checksum(&packet); + let translated_icmp = &packet[IPV4_MIN_HEADER_LEN..]; + assert_eq!( + read_u16(translated_icmp, ICMP_CHECKSUM_OFFSET), + original_icmp_checksum + ); + assert_eq!( + read_u16(translated_icmp, ICMP_CHECKSUM_OFFSET), + checksum_with_zeroed_word(translated_icmp, ICMP_CHECKSUM_OFFSET) + ); + } + + #[test] + fn rewrites_icmp_error_quoted_ipv4_and_visible_udp_checksum() { + let udp = udp_datagram(VIRTUAL_IP, REMOTE_IP, b"request", true); + let quoted = build_ipv4(VIRTUAL_IP, REMOTE_IP, IP_PROTOCOL_UDP, &udp, &[], 0); + let icmp = icmp_message(3, "ed); + let mut packet = build_ipv4(REMOTE_IP, VIRTUAL_IP, IP_PROTOCOL_ICMP, &icmp, &[], 0); + + rewrite_ipv4_destination(&mut packet, VIRTUAL_IP, CLIENT_IP).unwrap(); + + assert_valid_ipv4_checksum(&packet); + let outer_ihl = IPV4_MIN_HEADER_LEN; + let translated_icmp = &packet[outer_ihl..]; + assert_eq!( + read_u16(translated_icmp, ICMP_CHECKSUM_OFFSET), + checksum_with_zeroed_word(translated_icmp, ICMP_CHECKSUM_OFFSET) + ); + let translated_quote = &translated_icmp[ICMP_QUOTED_PACKET_OFFSET..]; + assert_eq!( + read_ipv4_address(translated_quote, IPV4_SOURCE_OFFSET), + CLIENT_IP + ); + assert_valid_ipv4_checksum(translated_quote); + assert_valid_transport_checksum(translated_quote, IP_PROTOCOL_UDP); + } + + #[test] + fn rewrites_icmp_error_quoted_destination_and_visible_tcp_checksum() { + let tcp = tcp_segment(REMOTE_IP, CLIENT_IP, b"request"); + let quoted = build_ipv4(REMOTE_IP, CLIENT_IP, IP_PROTOCOL_TCP, &tcp, &[], 0); + let icmp = icmp_message(11, "ed); + let mut packet = build_ipv4(CLIENT_IP, REMOTE_IP, IP_PROTOCOL_ICMP, &icmp, &[], 0); + + rewrite_ipv4_source(&mut packet, CLIENT_IP, VIRTUAL_IP).unwrap(); + + assert_valid_ipv4_checksum(&packet); + let translated_icmp = &packet[IPV4_MIN_HEADER_LEN..]; + assert_eq!( + read_u16(translated_icmp, ICMP_CHECKSUM_OFFSET), + checksum_with_zeroed_word(translated_icmp, ICMP_CHECKSUM_OFFSET) + ); + let translated_quote = &translated_icmp[ICMP_QUOTED_PACKET_OFFSET..]; + assert_eq!( + read_ipv4_address(translated_quote, IPV4_DESTINATION_OFFSET), + VIRTUAL_IP + ); + assert_valid_ipv4_checksum(translated_quote); + assert_valid_transport_checksum(translated_quote, IP_PROTOCOL_TCP); + } + + #[test] + fn rewrites_fragmented_tcp_checksum_only_in_first_fragment() { + let tcp = tcp_segment(CLIENT_IP, REMOTE_IP, b"0123456789abcdef01234567"); + let split = 24; + let mut first = build_ipv4( + CLIENT_IP, + REMOTE_IP, + IP_PROTOCOL_TCP, + &tcp[..split], + &[], + 0x2000, + ); + let mut second = build_ipv4( + CLIENT_IP, + REMOTE_IP, + IP_PROTOCOL_TCP, + &tcp[split..], + &[], + u16::try_from(split / 8).unwrap(), + ); + let second_payload_before = second[IPV4_MIN_HEADER_LEN..].to_vec(); + + rewrite_ipv4_source(&mut first, CLIENT_IP, VIRTUAL_IP).unwrap(); + rewrite_ipv4_source(&mut second, CLIENT_IP, VIRTUAL_IP).unwrap(); + + assert_valid_ipv4_checksum(&first); + assert_valid_ipv4_checksum(&second); + assert_eq!(&second[IPV4_MIN_HEADER_LEN..], second_payload_before); + let mut translated_tcp = first[IPV4_MIN_HEADER_LEN..].to_vec(); + translated_tcp.extend_from_slice(&second[IPV4_MIN_HEADER_LEN..]); + assert_eq!( + read_u16(&translated_tcp, TCP_CHECKSUM_OFFSET), + transport_checksum(VIRTUAL_IP, REMOTE_IP, IP_PROTOCOL_TCP, &translated_tcp) + ); + } + + #[test] + fn rejects_truncated_and_length_mismatched_ipv4_packets() { + let mut short = vec![0; IPV4_MIN_HEADER_LEN - 1]; + assert_eq!( + rewrite_ipv4_source(&mut short, CLIENT_IP, VIRTUAL_IP), + Err(Ipv4TranslationError::PacketTooShort { + actual: IPV4_MIN_HEADER_LEN - 1 + }) + ); + + let udp = udp_datagram(CLIENT_IP, REMOTE_IP, b"data", true); + let mut mismatched = build_ipv4(CLIENT_IP, REMOTE_IP, IP_PROTOCOL_UDP, &udp, &[], 0); + let declared = mismatched.len(); + mismatched.push(0); + assert_eq!( + rewrite_ipv4_source(&mut mismatched, CLIENT_IP, VIRTUAL_IP), + Err(Ipv4TranslationError::TotalLengthMismatch { + declared, + actual: declared + 1, + }) + ); + + let mut truncated_options = vec![0; IPV4_MIN_HEADER_LEN]; + truncated_options[0] = 0x46; + write_u16(&mut truncated_options, 2, IPV4_MIN_HEADER_LEN as u16); + assert_eq!( + rewrite_ipv4_source(&mut truncated_options, CLIENT_IP, VIRTUAL_IP), + Err(Ipv4TranslationError::TruncatedHeader { + header_len: 24, + actual: IPV4_MIN_HEADER_LEN, + }) + ); + } + + #[test] + fn rejects_unsupported_protocol_without_mutating_packet() { + let mut packet = build_ipv4(CLIENT_IP, REMOTE_IP, 47, &[0; 8], &[], 0); + let original = packet.clone(); + + assert_eq!( + rewrite_ipv4_source(&mut packet, CLIENT_IP, VIRTUAL_IP), + Err(Ipv4TranslationError::UnsupportedProtocol { protocol: 47 }) + ); + assert_eq!(packet, original); + } + + #[test] + fn rejects_unexpected_source_and_destination_without_mutating_packet() { + let tcp = tcp_segment(CLIENT_IP, REMOTE_IP, b"payload"); + let packet = build_ipv4(CLIENT_IP, REMOTE_IP, IP_PROTOCOL_TCP, &tcp, &[], 0); + + let mut source_packet = packet.clone(); + assert_eq!( + rewrite_ipv4_source(&mut source_packet, VIRTUAL_IP, CLIENT_IP), + Err(Ipv4TranslationError::UnexpectedSource { + expected: VIRTUAL_IP, + actual: CLIENT_IP, + }) + ); + assert_eq!(source_packet, packet); + + let mut destination_packet = packet.clone(); + assert_eq!( + rewrite_ipv4_destination(&mut destination_packet, VIRTUAL_IP, CLIENT_IP), + Err(Ipv4TranslationError::UnexpectedDestination { + expected: VIRTUAL_IP, + actual: REMOTE_IP, + }) + ); + assert_eq!(destination_packet, packet); + } + + #[test] + fn rejects_truncated_transport_and_icmp_quote() { + let mut tcp = build_ipv4( + CLIENT_IP, + REMOTE_IP, + IP_PROTOCOL_TCP, + &[0; TCP_MIN_HEADER_LEN - 1], + &[], + 0, + ); + assert_eq!( + rewrite_ipv4_source(&mut tcp, CLIENT_IP, VIRTUAL_IP), + Err(Ipv4TranslationError::TruncatedTransportHeader { + protocol: "TCP", + required: TCP_MIN_HEADER_LEN, + actual: TCP_MIN_HEADER_LEN - 1, + }) + ); + + let icmp = icmp_message(11, &[0; IPV4_MIN_HEADER_LEN - 1]); + let mut packet = build_ipv4(REMOTE_IP, VIRTUAL_IP, IP_PROTOCOL_ICMP, &icmp, &[], 0); + assert_eq!( + rewrite_ipv4_destination(&mut packet, VIRTUAL_IP, CLIENT_IP), + Err(Ipv4TranslationError::QuotedPacketTooShort { + actual: IPV4_MIN_HEADER_LEN - 1, + }) + ); + } +} diff --git a/easytier-core/src/gateway/vpn_portal/runtime.rs b/easytier-core/src/gateway/vpn_portal/runtime.rs new file mode 100644 index 00000000..ca5349e8 --- /dev/null +++ b/easytier-core/src/gateway/vpn_portal/runtime.rs @@ -0,0 +1,1655 @@ +//! Protocol-neutral portal session orchestration. +//! +//! Native adapters authenticate clients and yield sessions. This module owns +//! configured client identities, attached-peer lifetimes, per-client +//! generations, and IPv4 address translation at the Host packet seam. + +use std::{ + collections::{BTreeMap, BTreeSet}, + net::{IpAddr, Ipv4Addr}, + sync::Arc, +}; + +use async_trait::async_trait; +use cidr::Ipv4Inet; +use serde::{Deserialize, Serialize}; +use tokio::{ + sync::{Mutex, RwLock, mpsc, watch}, + task::{JoinHandle, JoinSet}, +}; +use tokio_util::sync::CancellationToken; + +use crate::{ + config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfigStore}, + events::{CoreEvent, CoreEventSink}, + peers::{ + attached::{AttachedPeerConfig, AttachedPeerRuntime}, + peer_manager::PeerManagerCore, + }, + socket::SocketListener, +}; + +use super::ipv4_translator::{rewrite_ipv4_destination, rewrite_ipv4_source}; + +pub const MAX_VPN_PORTAL_CLIENTS: usize = 64; +pub const DEFAULT_PORTAL_CLIENT_ADDRESS: Ipv4Addr = Ipv4Addr::new(192, 0, 2, 1); + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct PortalClientConfig { + pub name: String, + pub virtual_ip: Ipv4Addr, + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub groups: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct PortalRuntimeConfig { + pub clients: Vec, +} + +/// One authenticated protocol session produced by a native portal adapter. A +/// new value is emitted only for a new authenticated client generation; +/// ordinary reauthentication and endpoint roaming stay within that adapter. +/// The endpoint watch exposes the current authenticated endpoint and any later +/// roaming within the generation. Packet channels carry complete raw IPv4 +/// packets without protocol framing. +pub struct PortalSession { + pub client_name: String, + pub endpoint: watch::Receiver, + pub identity_private_key: [u8; 32], + pub from_client: mpsc::Receiver>, + pub to_client: mpsc::Sender>, +} + +impl std::fmt::Debug for PortalSession { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let endpoint = self.endpoint.borrow(); + formatter + .debug_struct("PortalSession") + .field("client_name", &self.client_name) + .field("endpoint", &endpoint.as_str()) + .finish_non_exhaustive() + } +} + +pub type PortalListener = Box>; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PortalClientConfigPlan { + pub name: String, + pub address: Ipv4Addr, + pub allowed_ips: Vec, + pub listener_url: url::Url, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PortalClientState { + Offline, + Connecting, + Online, + Error, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PortalClientInfoSnapshot { + pub name: String, + pub virtual_ip: Ipv4Addr, + pub groups: Vec, + pub state: PortalClientState, + pub peer_id: Option, + pub endpoint: Option, + pub tunnel_ip: Option, + pub client_config: String, + pub error: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PortalInfoSnapshot { + pub vpn_type: String, + pub clients: Vec, + pub listener: Option, +} + +#[async_trait] +pub trait PortalHost: Send + Sync + 'static { + /// Starts the shared protocol engines. Core owns accepted generations and + /// all portable attached-peer lifecycle after this seam. + async fn start_listeners(&self) -> anyhow::Result>; + + fn name(&self) -> String; + + fn render_client_config(&self, plan: &PortalClientConfigPlan) -> String; +} + +#[derive(Debug, Clone)] +struct ClientStatus { + state: PortalClientState, + generation: u64, + peer_id: Option, + endpoint: Option, + tunnel_ip: Option, + error: Option, +} + +impl Default for ClientStatus { + fn default() -> Self { + Self { + state: PortalClientState::Offline, + generation: 0, + peer_id: None, + endpoint: None, + tunnel_ip: None, + error: None, + } + } +} + +struct PortalRuntime { + cancel: CancellationToken, + tasks: JoinSet<()>, + listener_urls: Vec, +} + +pub struct PortalModule { + operation: Mutex<()>, + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + config: Option, + host: Option>, + events: Arc, + statuses: Arc>>, + session_locks: Arc>>>, + runtime: Mutex>, +} + +impl PortalModule { + pub fn new( + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + config: Option, + host: Option>, + events: Arc, + ) -> anyhow::Result> { + if let Some(config) = config.as_ref() { + validate_config(config, runtime_config.snapshot().as_ref())?; + } + let statuses = config + .as_ref() + .map(|config| { + config + .clients + .iter() + .map(|client| (client.name.clone(), ClientStatus::default())) + .collect() + }) + .unwrap_or_default(); + let session_locks = config + .as_ref() + .map(|config| { + config + .clients + .iter() + .map(|client| (client.name.clone(), Arc::new(Mutex::new(())))) + .collect() + }) + .unwrap_or_default(); + Ok(Arc::new(Self { + operation: Mutex::new(()), + peer_manager, + runtime_config, + config, + host, + events, + statuses: Arc::new(RwLock::new(statuses)), + session_locks: Arc::new(session_locks), + runtime: Mutex::new(None), + })) + } + #[cfg(feature = "vpn-portal")] + pub(crate) fn validate_runtime_config( + &self, + runtime_config: &CoreInstanceRuntimeConfig, + ) -> anyhow::Result<()> { + let Some(config) = self.config.as_ref() else { + return Ok(()); + }; + validate_runtime_compatibility(config, runtime_config) + } + + pub async fn start(&self) -> anyhow::Result<()> { + let _operation = self.operation.lock().await; + if self.config.is_none() { + return Ok(()); + } + let stale_runtime = { + let mut runtime = self.runtime.lock().await; + if runtime + .as_ref() + .is_some_and(|runtime| !runtime.cancel.is_cancelled()) + { + return Ok(()); + } + runtime.take() + }; + if let Some(stale_runtime) = stale_runtime { + self.shutdown_runtime(stale_runtime).await; + } + + let host = self.host.as_ref().ok_or_else(|| { + anyhow::anyhow!("VPN portal is configured but no host adapter exists") + })?; + let listeners = host.start_listeners().await?; + if listeners.is_empty() { + anyhow::bail!("VPN portal host returned no active listeners"); + } + + let mut prepared = Vec::with_capacity(listeners.len()); + for mut listener in listeners { + listener.listen().await?; + let listener_url = listener.local_url(); + prepared.push((listener, listener_url)); + } + + let cancel = CancellationToken::new(); + let start_signal = CancellationToken::new(); + let mut tasks = JoinSet::new(); + let mut listener_urls = Vec::with_capacity(prepared.len()); + for (listener, listener_url) in prepared { + listener_urls.push(listener_url.clone()); + tasks.spawn(Self::run_listener( + listener, + listener_url, + self.peer_manager.clone(), + self.runtime_config.clone(), + self.config.clone().expect("checked above"), + self.statuses.clone(), + self.session_locks.clone(), + self.events.clone(), + cancel.clone(), + start_signal.clone(), + )); + } + *self.runtime.lock().await = Some(PortalRuntime { + cancel, + tasks, + listener_urls: listener_urls.clone(), + }); + for local_url in &listener_urls { + self.events + .emit(CoreEvent::VpnPortalStarted(local_url.to_string())); + } + start_signal.cancel(); + Ok(()) + } + + #[allow(clippy::too_many_arguments)] + async fn run_listener( + mut listener: PortalListener, + listener_url: url::Url, + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + config: PortalRuntimeConfig, + statuses: Arc>>, + session_locks: Arc>>>, + events: Arc, + cancel: CancellationToken, + start_signal: CancellationToken, + ) { + tokio::select! { + _ = cancel.cancelled() => return, + _ = start_signal.cancelled() => {} + } + let mut sessions = JoinSet::new(); + loop { + tokio::select! { + _ = cancel.cancelled() => { + break; + } + accepted = listener.accept() => match accepted { + Ok(session) => { + sessions.spawn(Self::run_session( + session, + listener_url.clone(), + peer_manager.clone(), + runtime_config.clone(), + config.clone(), + statuses.clone(), + session_locks.clone(), + events.clone(), + cancel.clone(), + )); + } + Err(error) => { + tracing::warn!(?error, "VPN portal listener stopped accepting"); + cancel.cancel(); + break; + } + }, + _ = sessions.join_next(), if !sessions.is_empty() => {} + } + } + drop(listener); + while sessions.join_next().await.is_some() {} + } + + #[allow(clippy::too_many_arguments)] + async fn run_session( + mut session: PortalSession, + listener_url: url::Url, + peer_manager: Arc, + runtime_config: CoreRuntimeConfigStore, + config: PortalRuntimeConfig, + statuses: Arc>>, + session_locks: Arc>>>, + events: Arc, + cancel: CancellationToken, + ) { + let Some(client) = config + .clients + .iter() + .find(|client| client.name == session.client_name) + .cloned() + else { + tracing::warn!(client = %session.client_name, "unknown VPN portal client session"); + return; + }; + let session_lock = session_locks + .get(&client.name) + .expect("validated client session lock exists"); + let _session_guard = tokio::select! { + _ = cancel.cancelled() => return, + guard = session_lock.lock() => guard, + }; + let generation = { + let mut statuses = statuses.write().await; + let status = statuses + .get_mut(&client.name) + .expect("validated client status exists"); + status.generation = status.generation.wrapping_add(1); + status.state = PortalClientState::Connecting; + status.endpoint = Some(session.endpoint.borrow_and_update().clone()); + status.tunnel_ip = None; + status.error = None; + status.generation + }; + + let attached = match AttachedPeerRuntime::connect( + peer_manager, + runtime_config, + AttachedPeerConfig { + name: client.name.clone(), + virtual_ip: client.virtual_ip, + groups: client.groups.clone(), + identity_private_key: session.identity_private_key, + }, + ) + .await + { + Ok(attached) => attached, + Err(error) => { + Self::finish_generation( + &statuses, + &client.name, + generation, + Some(error.to_string()), + ) + .await; + return; + } + }; + + { + let mut statuses = statuses.write().await; + let status = statuses + .get_mut(&client.name) + .expect("validated client status exists"); + if status.generation != generation { + drop(statuses); + attached.close().await; + return; + } + status.state = PortalClientState::Online; + status.peer_id = Some(attached.peer_id()); + } + events.emit(CoreEvent::VpnPortalClientConnected { + portal: listener_url.to_string(), + client: client.name.clone(), + }); + + let mut client_stream = session.from_client; + let endpoint = session.endpoint; + let client_sink = session.to_client; + let client_ip = Arc::new(Mutex::new(None::)); + let client_to_mesh = { + let attached = attached.clone(); + let client_ip = client_ip.clone(); + let statuses = statuses.clone(); + let name = client.name.clone(); + let virtual_ip = client.virtual_ip; + tokio::spawn(async move { + while let Some(mut payload) = client_stream.recv().await { + let Some(source) = ipv4_source(&payload) else { + continue; + }; + match *client_ip.lock().await { + Some(expected) if expected != source => { + tracing::warn!(client = %name, ?expected, ?source, "VPN client source changed"); + continue; + } + None | Some(_) => {} + } + if rewrite_ipv4_source(&mut payload, source, virtual_ip).is_err() { + continue; + } + let learned = { + let mut tunnel_ip = client_ip.lock().await; + if tunnel_ip.is_none() { + *tunnel_ip = Some(source); + true + } else { + false + } + }; + if learned { + let mut statuses = statuses.write().await; + if let Some(status) = statuses.get_mut(&name) + && status.generation == generation + { + status.tunnel_ip = Some(source); + } + } + if let Err(error) = attached.send_packet(&payload).await { + tracing::debug!(?error, client = %name, "attached peer send failed"); + break; + } + } + }) + }; + let mesh_to_client = { + let attached = attached.clone(); + let client_ip = client_ip.clone(); + let statuses = statuses.clone(); + let name = client.name.clone(); + let virtual_ip = client.virtual_ip; + tokio::spawn(async move { + while let Some(packet) = attached.recv_packet().await { + let Some(tunnel_ip) = *client_ip.lock().await else { + continue; + }; + let mut payload = packet.payload().to_vec(); + if rewrite_ipv4_destination(&mut payload, virtual_ip, tunnel_ip).is_err() { + continue; + } + if client_sink.send(payload).await.is_err() { + break; + } + let mut statuses = statuses.write().await; + if let Some(status) = statuses.get_mut(&name) + && status.generation == generation + { + status.tunnel_ip = Some(tunnel_ip); + } + } + }) + }; + Self::supervise_session_io( + endpoint, + statuses.clone(), + client.name.clone(), + generation, + cancel, + client_to_mesh, + mesh_to_client, + ) + .await; + + attached.close().await; + Self::finish_generation(&statuses, &client.name, generation, None).await; + events.emit(CoreEvent::VpnPortalClientDisconnected { + portal: listener_url.to_string(), + client: client.name, + }); + } + #[allow(clippy::too_many_arguments)] + async fn supervise_session_io( + mut endpoint: watch::Receiver, + statuses: Arc>>, + client_name: String, + generation: u64, + cancel: CancellationToken, + mut client_to_mesh: JoinHandle<()>, + mut mesh_to_client: JoinHandle<()>, + ) { + let mut client_to_mesh_finished = false; + let mut mesh_to_client_finished = false; + loop { + tokio::select! { + biased; + _ = cancel.cancelled() => break, + result = &mut client_to_mesh, if !client_to_mesh_finished => { + client_to_mesh_finished = true; + if let Err(error) = result { + tracing::debug!( + ?error, + client = %client_name, + "VPN portal client-to-mesh task failed" + ); + } + break; + } + result = &mut mesh_to_client, if !mesh_to_client_finished => { + mesh_to_client_finished = true; + if let Err(error) = result { + tracing::debug!( + ?error, + client = %client_name, + "VPN portal mesh-to-client task failed" + ); + } + break; + } + changed = endpoint.changed() => { + if changed.is_err() { + break; + } + let current = endpoint.borrow_and_update().clone(); + let mut statuses = statuses.write().await; + if let Some(status) = statuses.get_mut(&client_name) + && status.generation == generation + { + status.endpoint = Some(current); + } + } + } + } + if !client_to_mesh_finished { + client_to_mesh.abort(); + let _ = client_to_mesh.await; + } + if !mesh_to_client_finished { + mesh_to_client.abort(); + let _ = mesh_to_client.await; + } + } + + async fn finish_generation( + statuses: &RwLock>, + name: &str, + generation: u64, + error: Option, + ) { + let mut statuses = statuses.write().await; + let Some(status) = statuses.get_mut(name) else { + return; + }; + if status.generation != generation { + return; + } + status.state = if error.is_some() { + PortalClientState::Error + } else { + PortalClientState::Offline + }; + status.peer_id = None; + status.endpoint = None; + status.tunnel_ip = None; + status.error = error; + } + + async fn shutdown_runtime(&self, mut runtime: PortalRuntime) { + runtime.cancel.cancel(); + while let Some(result) = runtime.tasks.join_next().await { + if let Err(error) = result { + tracing::debug!(?error, "VPN portal listener task failed"); + } + } + for status in self.statuses.write().await.values_mut() { + *status = ClientStatus::default(); + } + } + + pub async fn stop(&self) { + let _operation = self.operation.lock().await; + let Some(runtime) = self.runtime.lock().await.take() else { + return; + }; + self.shutdown_runtime(runtime).await; + } + + pub async fn info_snapshot(&self) -> PortalInfoSnapshot { + let Some(config) = self.config.as_ref() else { + return PortalInfoSnapshot { + vpn_type: "null".to_owned(), + clients: Vec::new(), + listener: None, + }; + }; + let listener_url = self + .runtime + .lock() + .await + .as_ref() + .filter(|runtime| !runtime.cancel.is_cancelled()) + .and_then(|runtime| runtime.listener_urls.first().cloned()); + let allowed_ips = self.allowed_ips().await; + let statuses = self.statuses.read().await; + let clients = config + .clients + .iter() + .map(|client| { + let status = statuses.get(&client.name).cloned().unwrap_or_default(); + let client_config = match (self.host.as_ref(), listener_url.as_ref()) { + (Some(host), Some(listener_url)) => { + host.render_client_config(&PortalClientConfigPlan { + name: client.name.clone(), + address: DEFAULT_PORTAL_CLIENT_ADDRESS, + allowed_ips: allowed_ips.clone(), + listener_url: listener_url.clone(), + }) + } + _ => String::new(), + }; + PortalClientInfoSnapshot { + name: client.name.clone(), + virtual_ip: client.virtual_ip, + groups: client.groups.clone(), + state: status.state, + peer_id: status.peer_id, + endpoint: status.endpoint, + tunnel_ip: status.tunnel_ip, + client_config, + error: status.error, + } + }) + .collect(); + PortalInfoSnapshot { + vpn_type: self + .host + .as_ref() + .map_or_else(|| "null".to_owned(), |host| host.name()), + clients, + listener: listener_url.map(|url| url.to_string()), + } + } + + async fn allowed_ips(&self) -> Vec { + let snapshot = self.runtime_config.snapshot(); + let mut allowed = BTreeSet::new(); + for route in self.peer_manager.list_route_snapshots().await { + allowed.extend(route.proxy_cidrs); + } + if let Some(ipv4) = snapshot.peer.runtime.core.routes.ipv4.as_ref() + && let IpAddr::V4(address) = ipv4.address + && let Ok(inet) = Ipv4Inet::new(address, ipv4.prefix_len) + { + allowed.insert(inet.network().to_string()); + } + for proxy in &snapshot.peer.runtime.core.routes.proxy_networks { + let mapped = proxy.mapped.as_ref().unwrap_or(&proxy.real); + allowed.insert(format!("{}/{}", mapped.address, mapped.prefix_len)); + } + allowed.into_iter().collect() + } +} + +fn validate_config( + config: &PortalRuntimeConfig, + runtime_config: &CoreInstanceRuntimeConfig, +) -> anyhow::Result<()> { + if config.clients.is_empty() { + anyhow::bail!("VPN portal requires at least one configured client"); + } + if config.clients.len() > MAX_VPN_PORTAL_CLIENTS { + anyhow::bail!("VPN portal supports at most {MAX_VPN_PORTAL_CLIENTS} clients"); + } + validate_runtime_compatibility(config, runtime_config)?; + let snapshot = runtime_config; + let declared_groups = snapshot + .peer + .acl_group_declarations + .iter() + .map(|group| group.group_name.as_str()) + .collect::>(); + let mut names = BTreeSet::new(); + let mut addresses = BTreeSet::new(); + for client in &config.clients { + validate_client_name(&client.name)?; + if !names.insert(client.name.as_str()) { + anyhow::bail!("duplicate VPN portal client name: {}", client.name); + } + if !addresses.insert(client.virtual_ip) { + anyhow::bail!("duplicate VPN portal virtual IP: {}", client.virtual_ip); + } + for group in &client.groups { + if !declared_groups.contains(group.as_str()) { + anyhow::bail!( + "VPN portal client {} uses unknown ACL group {group}", + client.name + ); + } + } + } + Ok(()) +} + +fn validate_runtime_compatibility( + config: &PortalRuntimeConfig, + runtime_config: &CoreInstanceRuntimeConfig, +) -> anyhow::Result<()> { + let snapshot = runtime_config; + if snapshot + .peer + .runtime + .network_identity + .network_secret + .as_deref() + .is_none_or(str::is_empty) + { + anyhow::bail!("VPN portal requires an admin node with a non-empty network secret"); + } + if snapshot.services.dhcp_ipv4 { + anyhow::bail!("VPN portal does not support DHCP IPv4 on the portal node"); + } + let prefix = snapshot + .peer + .runtime + .core + .routes + .ipv4 + .as_ref() + .ok_or_else(|| anyhow::anyhow!("VPN portal requires a static IPv4 address"))?; + let IpAddr::V4(portal_ip) = prefix.address else { + anyhow::bail!("VPN portal requires an IPv4 route prefix"); + }; + let network = Ipv4Inet::new(portal_ip, prefix.prefix_len) + .map_err(|error| anyhow::anyhow!("invalid portal IPv4 prefix: {error}"))? + .network(); + for client in &config.clients { + if client.virtual_ip == portal_ip + || !network.contains(&client.virtual_ip) + || client.virtual_ip == network.first_address() + || client.virtual_ip == network.last_address() + { + anyhow::bail!( + "VPN portal client {} has an unusable virtual IP {}", + client.name, + client.virtual_ip + ); + } + } + Ok(()) +} + +fn validate_client_name(name: &str) -> anyhow::Result<()> { + let valid = !name.is_empty() + && name.len() <= 63 + && name + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-') + && name + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphanumeric) + && name + .as_bytes() + .last() + .is_some_and(u8::is_ascii_alphanumeric); + if !valid { + anyhow::bail!("invalid VPN portal client name: {name}"); + } + Ok(()) +} + +fn ipv4_source(payload: &[u8]) -> Option { + if payload.len() < 20 || payload[0] >> 4 != 4 { + return None; + } + Some(Ipv4Addr::new( + payload[12], + payload[13], + payload[14], + payload[15], + )) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + config::{ + IpPrefix, NetworkIdentity, + peers::{PeerGroupIdentity, PeerRuntimeSnapshot}, + runtime::CoreRuntimeConfig, + }, + peers::peer_manager::PeerManagerCore, + }; + use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; + use std::{ + future::pending, + sync::{ + Mutex as StdMutex, Weak, + atomic::{AtomicBool, AtomicUsize, Ordering}, + }, + }; + use tokio::sync::{Notify, mpsc}; + use x25519_dalek::{PublicKey, StaticSecret}; + + #[derive(Default)] + struct RecordingEvents(std::sync::Mutex>); + + impl CoreEventSink for RecordingEvents { + fn emit(&self, event: CoreEvent) { + self.0.lock().unwrap().push(event); + } + } + + struct StaticPortalHost { + listeners: StdMutex>>, + } + + impl StaticPortalHost { + fn new(listeners: Vec) -> Arc { + Arc::new(Self { + listeners: StdMutex::new(Some(listeners)), + }) + } + } + + #[async_trait] + impl PortalHost for StaticPortalHost { + async fn start_listeners(&self) -> anyhow::Result> { + self.listeners + .lock() + .unwrap() + .take() + .ok_or_else(|| anyhow::anyhow!("test listeners already started")) + } + + fn name(&self) -> String { + "test".to_owned() + } + + fn render_client_config(&self, plan: &PortalClientConfigPlan) -> String { + format!("config:{}", plan.name) + } + } + + #[derive(Debug)] + struct PendingPortalListener { + url: url::Url, + accept_calls: Arc, + } + + #[async_trait] + impl SocketListener for PendingPortalListener { + type Accepted = PortalSession; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + self.accept_calls.fetch_add(1, Ordering::SeqCst); + pending().await + } + + fn local_url(&self) -> url::Url { + self.url.clone() + } + } + + #[derive(Debug)] + struct FailingListenPortalListener { + url: url::Url, + listening: Arc, + fail: Arc, + } + + #[async_trait] + impl SocketListener for FailingListenPortalListener { + type Accepted = PortalSession; + + async fn listen(&mut self) -> anyhow::Result<()> { + self.listening.notify_one(); + self.fail.notified().await; + anyhow::bail!("listener setup failed") + } + + async fn accept(&mut self) -> anyhow::Result { + unreachable!("failed listeners cannot accept") + } + + fn local_url(&self) -> url::Url { + self.url.clone() + } + } + + #[derive(Debug)] + struct FailingAcceptPortalListener { + url: url::Url, + } + + #[async_trait] + impl SocketListener for FailingAcceptPortalListener { + type Accepted = PortalSession; + + async fn listen(&mut self) -> anyhow::Result<()> { + Ok(()) + } + + async fn accept(&mut self) -> anyhow::Result { + anyhow::bail!("listener receive failed") + } + + fn local_url(&self) -> url::Url { + self.url.clone() + } + } + + struct RestartingPortalHost { + starts: AtomicUsize, + accept_calls: Arc, + } + + #[async_trait] + impl PortalHost for RestartingPortalHost { + async fn start_listeners(&self) -> anyhow::Result> { + let attempt = self.starts.fetch_add(1, Ordering::SeqCst); + let url = format!("test://127.0.0.1:{}", 10000 + attempt) + .parse() + .unwrap(); + if attempt == 0 { + Ok(vec![Box::new(FailingAcceptPortalListener { url })]) + } else { + Ok(vec![Box::new(PendingPortalListener { + url, + accept_calls: self.accept_calls.clone(), + })]) + } + } + + fn name(&self) -> String { + "test".to_owned() + } + + fn render_client_config(&self, plan: &PortalClientConfigPlan) -> String { + format!("config:{}", plan.name) + } + } + + #[derive(Default)] + struct RuntimeObservingEvents { + module: StdMutex>, + runtime_visible_on_start: AtomicBool, + } + + impl CoreEventSink for RuntimeObservingEvents { + fn emit(&self, event: CoreEvent) { + if matches!(event, CoreEvent::VpnPortalStarted(_)) + && let Some(module) = self.module.lock().unwrap().upgrade() + && let Ok(runtime) = module.runtime.try_lock() + { + self.runtime_visible_on_start.store( + runtime + .as_ref() + .is_some_and(|runtime| !runtime.cancel.is_cancelled()), + Ordering::SeqCst, + ); + } + } + } + + struct TaskDropCounter(Arc); + + impl Drop for TaskDropCounter { + fn drop(&mut self) { + self.0.fetch_add(1, Ordering::SeqCst); + } + } + + fn runtime_config() -> CoreRuntimeConfigStore { + let mut peer = PeerRuntimeSnapshot::default(); + peer.runtime.network_identity = + NetworkIdentity::new("portal-test".to_owned(), "shared-secret".to_owned()); + peer.runtime.core.routes.ipv4 = Some(IpPrefix { + address: IpAddr::V4(Ipv4Addr::new(10, 82, 0, 1)), + prefix_len: 24, + }); + peer.acl_group_declarations = vec![PeerGroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }]; + CoreRuntimeConfigStore::new(CoreRuntimeConfig::default(), Arc::new(peer)) + } + + fn client(name: &str, virtual_ip: Ipv4Addr, groups: &[&str]) -> PortalClientConfig { + PortalClientConfig { + name: name.to_owned(), + virtual_ip, + groups: groups.iter().map(|group| (*group).to_owned()).collect(), + } + } + + fn raw_ipv4(source: Ipv4Addr, destination: Ipv4Addr) -> Vec { + let mut packet = vec![0u8; 28]; + packet[0] = 0x45; + packet[2..4].copy_from_slice(&28u16.to_be_bytes()); + packet[8] = 64; + packet[9] = 1; + packet[12..16].copy_from_slice(&source.octets()); + packet[16..20].copy_from_slice(&destination.octets()); + packet[20] = 8; + packet + } + fn network_runtime() -> (Arc, CoreRuntimeConfigStore) { + network_runtime_with_secure_mode(false) + } + + fn secure_network_runtime() -> (Arc, CoreRuntimeConfigStore) { + network_runtime_with_secure_mode(true) + } + + fn network_runtime_with_secure_mode( + secure_mode: bool, + ) -> (Arc, CoreRuntimeConfigStore) { + let store = runtime_config(); + if secure_mode { + let private = StaticSecret::from([42; 32]); + let public = PublicKey::from(&private); + store.update_peer_with(|peer| { + peer.runtime.secure_mode = Some(crate::proto::common::SecureModeConfig { + enabled: true, + local_private_key: Some(BASE64_STANDARD.encode(private.to_bytes())), + local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), + }); + }); + } + let snapshot = store.snapshot(); + let runtime = snapshot.peer.runtime.clone(); + let mut portable = crate::peers::peer_manager::PortablePeerManagerConfig::new(runtime); + portable.snapshot.acl_group_declarations = snapshot.peer.acl_group_declarations.clone(); + let (packet_sender, _packet_receiver) = crate::host::packet::host_packet_channel(); + let public_ipv6_runtime = crate::peers::public_ipv6::CorePublicIpv6Runtime::new( + store.clone(), + Arc::new(()), + Arc::new(()), + ); + let peer = Arc::new( + PeerManagerCore::new( + portable, + store.clone(), + Arc::new(()), + packet_sender, + public_ipv6_runtime, + Arc::new(()), + None, + Arc::new(()), + ) + .unwrap(), + ); + (peer, store) + } + + #[test] + fn portal_runtime_rejects_duplicate_names_and_virtual_ips() { + let runtime_config = runtime_config(); + let snapshot = runtime_config.snapshot(); + let alice = client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"]); + + let duplicate_name = PortalRuntimeConfig { + clients: vec![ + alice.clone(), + client("alice", Ipv4Addr::new(10, 82, 0, 3), &["ops"]), + ], + }; + let error = validate_config(&duplicate_name, snapshot.as_ref()) + .unwrap_err() + .to_string(); + assert!(error.contains("duplicate VPN portal client name")); + + let duplicate_ip = PortalRuntimeConfig { + clients: vec![alice, client("bob", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }; + let error = validate_config(&duplicate_ip, snapshot.as_ref()) + .unwrap_err() + .to_string(); + assert!(error.contains("duplicate VPN portal virtual IP")); + } + + #[test] + fn portal_runtime_rejects_unknown_acl_groups() { + let runtime_config = runtime_config(); + let config = PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["unknown"])], + }; + let error = validate_config(&config, runtime_config.snapshot().as_ref()) + .unwrap_err() + .to_string(); + assert!(error.contains("unknown ACL group")); + } + + #[test] + fn portal_runtime_rejects_credential_node() { + let runtime_config = runtime_config(); + runtime_config.update_peer_with(|peer| { + peer.runtime.network_identity.network_secret = None; + }); + let config = PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }; + + let error = validate_config(&config, runtime_config.snapshot().as_ref()) + .unwrap_err() + .to_string(); + + assert!(error.contains("VPN portal requires an admin node")); + } + + #[test] + fn portal_session_debug_redacts_identity_private_key() { + let identity_private_key = [173u8; 32]; + let (_to_runtime, from_client) = mpsc::channel(1); + let (to_client, _from_runtime) = mpsc::channel(1); + let (_endpoint_sender, endpoint) = tokio::sync::watch::channel("portal://alice".to_owned()); + let session = PortalSession { + client_name: "alice".to_owned(), + endpoint, + identity_private_key, + from_client, + to_client, + }; + + let rendered = format!("{session:?}"); + assert!(!rendered.contains("identity_private_key")); + assert!(!rendered.contains(&format!("{identity_private_key:?}"))); + } + + #[tokio::test] + async fn portal_session_publishes_learned_tunnel_ip_without_mesh_reply() { + let (peer_manager, runtime_config) = network_runtime(); + peer_manager.run().await.unwrap(); + let config = PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }; + let statuses = Arc::new(RwLock::new(BTreeMap::from([( + "alice".to_owned(), + ClientStatus::default(), + )]))); + let session_locks = Arc::new(BTreeMap::from([( + "alice".to_owned(), + Arc::new(Mutex::new(())), + )])); + let (to_runtime, from_client) = mpsc::channel(1); + let (to_client, _from_runtime) = mpsc::channel(1); + let (endpoint_sender, endpoint) = tokio::sync::watch::channel("portal://alice".to_owned()); + let session = PortalSession { + client_name: "alice".to_owned(), + endpoint, + identity_private_key: [173u8; 32], + from_client, + to_client, + }; + let cancel = CancellationToken::new(); + let task = tokio::spawn(PortalModule::run_session( + session, + "portal://listener".parse().unwrap(), + peer_manager.clone(), + runtime_config, + config, + statuses.clone(), + session_locks, + Arc::new(()), + cancel, + )); + + to_runtime + .send(raw_ipv4( + DEFAULT_PORTAL_CLIENT_ADDRESS, + Ipv4Addr::new(10, 82, 0, 1), + )) + .await + .unwrap(); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let status = statuses.read().await.get("alice").cloned().unwrap(); + if status.tunnel_ip == Some(DEFAULT_PORTAL_CLIENT_ADDRESS) { + return; + } + assert_ne!( + status.state, + PortalClientState::Error, + "portal session failed before learning tunnel IP: {:?}", + status.error + ); + assert!( + !task.is_finished(), + "portal session ended before learning tunnel IP" + ); + tokio::task::yield_now().await; + } + }) + .await + .expect("learned tunnel IP was not published"); + endpoint_sender.send("portal://roamed".to_owned()).unwrap(); + tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let status = statuses.read().await.get("alice").cloned().unwrap(); + if status.endpoint.as_deref() == Some("portal://roamed") { + return; + } + assert!( + !task.is_finished(), + "portal session ended before publishing the roamed endpoint" + ); + tokio::task::yield_now().await; + } + }) + .await + .expect("roamed endpoint was not published"); + + drop(to_runtime); + task.await.unwrap(); + peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn portal_session_cancellation_runs_complete_cleanup() { + let (peer_manager, runtime_config) = secure_network_runtime(); + peer_manager.run().await.unwrap(); + let config = PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }; + let statuses = Arc::new(RwLock::new(BTreeMap::from([( + "alice".to_owned(), + ClientStatus::default(), + )]))); + let session_locks = Arc::new(BTreeMap::from([( + "alice".to_owned(), + Arc::new(Mutex::new(())), + )])); + let (_to_runtime, from_client) = mpsc::channel(1); + let (to_client, _from_runtime) = mpsc::channel(1); + let (_endpoint_sender, endpoint) = tokio::sync::watch::channel("portal://alice".to_owned()); + let identity_private_key = [174u8; 32]; + let identity_public_key = + *PublicKey::from(&StaticSecret::from(identity_private_key)).as_bytes(); + let session = PortalSession { + client_name: "alice".to_owned(), + endpoint, + identity_private_key, + from_client, + to_client, + }; + let cancel = CancellationToken::new(); + let events = Arc::new(RecordingEvents::default()); + let mut task = tokio::spawn(PortalModule::run_session( + session, + "portal://listener".parse().unwrap(), + peer_manager.clone(), + runtime_config, + config, + statuses.clone(), + session_locks, + events.clone(), + cancel.clone(), + )); + + let attached_peer_id = tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let status = statuses.read().await.get("alice").cloned().unwrap(); + if status.state == PortalClientState::Online { + return status.peer_id.unwrap(); + } + assert!( + !task.is_finished(), + "portal session ended before becoming online: {:?}", + status.error + ); + tokio::task::yield_now().await; + } + }) + .await + .expect("portal session did not become online"); + assert!( + peer_manager + .credential_manager() + .is_pubkey_trusted(&identity_public_key) + ); + + cancel.cancel(); + if tokio::time::timeout(std::time::Duration::from_secs(5), &mut task) + .await + .is_err() + { + task.abort(); + let _ = task.await; + peer_manager.clear_resources().await; + panic!("portal session did not stop after cancellation"); + } + + let status = statuses.read().await.get("alice").cloned().unwrap(); + assert_eq!(status.state, PortalClientState::Offline); + assert!(status.peer_id.is_none()); + assert!(status.tunnel_ip.is_none()); + assert!( + !peer_manager + .get_peer_map() + .has_direct_attached_peer(attached_peer_id) + ); + assert!( + !peer_manager + .credential_manager() + .is_pubkey_trusted(&identity_public_key) + ); + { + let events = events.0.lock().unwrap(); + assert_eq!( + events + .iter() + .filter(|event| matches!(event, CoreEvent::VpnPortalClientConnected { .. })) + .count(), + 1 + ); + assert_eq!( + events + .iter() + .filter(|event| matches!(event, CoreEvent::VpnPortalClientDisconnected { .. })) + .count(), + 1 + ); + } + peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn portal_session_stops_when_outbound_packet_task_ends() { + let (peer_manager, runtime_config) = network_runtime(); + peer_manager.run().await.unwrap(); + let virtual_ip = Ipv4Addr::new(10, 82, 0, 2); + let config = PortalRuntimeConfig { + clients: vec![client("alice", virtual_ip, &["ops"])], + }; + let statuses = Arc::new(RwLock::new(BTreeMap::from([( + "alice".to_owned(), + ClientStatus::default(), + )]))); + let session_locks = Arc::new(BTreeMap::from([( + "alice".to_owned(), + Arc::new(Mutex::new(())), + )])); + let (to_runtime, from_client) = mpsc::channel(1); + let (to_client, from_runtime) = mpsc::channel(1); + let (_endpoint_sender, endpoint) = tokio::sync::watch::channel("portal://alice".to_owned()); + let session = PortalSession { + client_name: "alice".to_owned(), + endpoint, + identity_private_key: [175u8; 32], + from_client, + to_client, + }; + let events = Arc::new(RecordingEvents::default()); + let task = tokio::spawn(PortalModule::run_session( + session, + "portal://listener".parse().unwrap(), + peer_manager.clone(), + runtime_config, + config, + statuses.clone(), + session_locks, + events.clone(), + CancellationToken::new(), + )); + + to_runtime + .send(raw_ipv4( + DEFAULT_PORTAL_CLIENT_ADDRESS, + Ipv4Addr::new(10, 82, 0, 1), + )) + .await + .unwrap(); + let attached_peer_id = tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let status = statuses.read().await.get("alice").cloned().unwrap(); + if status.state == PortalClientState::Online + && status.tunnel_ip == Some(DEFAULT_PORTAL_CLIENT_ADDRESS) + { + return status.peer_id.unwrap(); + } + assert!( + !task.is_finished(), + "portal session ended before learning its tunnel address: {:?}", + status.error + ); + tokio::task::yield_now().await; + } + }) + .await + .expect("portal session did not learn its tunnel address"); + + drop(from_runtime); + let mesh_packet = raw_ipv4(Ipv4Addr::new(10, 82, 0, 1), virtual_ip); + if tokio::time::timeout(std::time::Duration::from_secs(5), async { + while !task.is_finished() { + let _ = peer_manager + .send_msg_by_ip( + crate::packet::ZCPacket::new_with_payload(&mesh_packet), + IpAddr::V4(virtual_ip), + false, + ) + .await; + tokio::task::yield_now().await; + } + }) + .await + .is_err() + { + task.abort(); + let _ = task.await; + peer_manager.clear_resources().await; + panic!("portal session did not stop when its outbound task ended"); + } + task.await.unwrap(); + + let status = statuses.read().await.get("alice").cloned().unwrap(); + assert_eq!(status.state, PortalClientState::Offline); + assert!(status.peer_id.is_none()); + assert!( + !peer_manager + .get_peer_map() + .has_direct_attached_peer(attached_peer_id) + ); + assert_eq!( + events + .0 + .lock() + .unwrap() + .iter() + .filter(|event| matches!(event, CoreEvent::VpnPortalClientDisconnected { .. })) + .count(), + 1 + ); + peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn portal_start_does_not_accept_before_all_listeners_are_ready() { + let (peer_manager, runtime_config) = network_runtime(); + let first_accept_calls = Arc::new(AtomicUsize::new(0)); + let second_listening = Arc::new(Notify::new()); + let fail_second = Arc::new(Notify::new()); + let host = StaticPortalHost::new(vec![ + Box::new(PendingPortalListener { + url: "test://127.0.0.1:10001".parse().unwrap(), + accept_calls: first_accept_calls.clone(), + }), + Box::new(FailingListenPortalListener { + url: "test://127.0.0.1:10002".parse().unwrap(), + listening: second_listening.clone(), + fail: fail_second.clone(), + }), + ]); + let events = Arc::new(RecordingEvents::default()); + let module = PortalModule::new( + peer_manager.clone(), + runtime_config, + Some(PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }), + Some(host), + events.clone(), + ) + .unwrap(); + let start = tokio::spawn({ + let module = module.clone(); + async move { module.start().await } + }); + second_listening.notified().await; + + let accepted_before_failure = + tokio::time::timeout(std::time::Duration::from_millis(50), async { + while first_accept_calls.load(Ordering::SeqCst) == 0 { + tokio::task::yield_now().await; + } + }) + .await + .is_ok(); + fail_second.notify_one(); + + let error = start.await.unwrap().unwrap_err(); + assert!(error.to_string().contains("listener setup failed")); + assert!( + !accepted_before_failure, + "an earlier listener accepted before startup completed" + ); + assert!(module.runtime.lock().await.is_none()); + assert!( + events + .0 + .lock() + .unwrap() + .iter() + .all(|event| !matches!(event, CoreEvent::VpnPortalStarted(_))) + ); + peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn portal_started_event_observes_installed_runtime() { + let (peer_manager, runtime_config) = network_runtime(); + let events = Arc::new(RuntimeObservingEvents::default()); + let module = PortalModule::new( + peer_manager.clone(), + runtime_config, + Some(PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }), + Some(StaticPortalHost::new(vec![Box::new( + PendingPortalListener { + url: "test://127.0.0.1:10003".parse().unwrap(), + accept_calls: Arc::new(AtomicUsize::new(0)), + }, + )])), + events.clone(), + ) + .unwrap(); + *events.module.lock().unwrap() = Arc::downgrade(&module); + + module.start().await.unwrap(); + + assert!( + events.runtime_visible_on_start.load(Ordering::SeqCst), + "VpnPortalStarted was emitted before runtime installation" + ); + module.stop().await; + peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn portal_restarts_after_listener_accept_failure() { + let (peer_manager, runtime_config) = network_runtime(); + let host = Arc::new(RestartingPortalHost { + starts: AtomicUsize::new(0), + accept_calls: Arc::new(AtomicUsize::new(0)), + }); + let module = PortalModule::new( + peer_manager.clone(), + runtime_config, + Some(PortalRuntimeConfig { + clients: vec![client("alice", Ipv4Addr::new(10, 82, 0, 2), &["ops"])], + }), + Some(host.clone()), + Arc::new(()), + ) + .unwrap(); + module.start().await.unwrap(); + + tokio::time::timeout(std::time::Duration::from_secs(1), async { + loop { + if module.info_snapshot().await.listener.is_none() { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("failed listener remained reported as active"); + + module.start().await.unwrap(); + + assert_eq!(host.starts.load(Ordering::SeqCst), 2); + assert_eq!( + module.info_snapshot().await.listener.as_deref(), + Some("test://127.0.0.1:10001") + ); + module.stop().await; + peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn session_io_cancellation_aborts_blocked_packet_directions() { + let drops = Arc::new(AtomicUsize::new(0)); + let blocked_task = |drops: Arc| { + tokio::spawn(async move { + let _drop_counter = TaskDropCounter(drops); + pending::<()>().await; + }) + }; + let client_to_mesh = blocked_task(drops.clone()); + let mesh_to_client = blocked_task(drops.clone()); + let (_endpoint_sender, endpoint) = watch::channel("test://127.0.0.1:10004".to_owned()); + let statuses = Arc::new(RwLock::new(BTreeMap::from([( + "alice".to_owned(), + ClientStatus { + generation: 1, + ..Default::default() + }, + )]))); + let cancel = CancellationToken::new(); + let supervisor = tokio::spawn(PortalModule::supervise_session_io( + endpoint, + statuses, + "alice".to_owned(), + 1, + cancel.clone(), + client_to_mesh, + mesh_to_client, + )); + tokio::task::yield_now().await; + + cancel.cancel(); + + tokio::time::timeout(std::time::Duration::from_secs(1), supervisor) + .await + .expect("session I/O supervisor ignored cancellation") + .unwrap(); + assert_eq!(drops.load(Ordering::SeqCst), 2); + } + + #[test] + fn client_names_are_dns_safe() { + assert!(validate_client_name("laptop-1").is_ok()); + assert!(validate_client_name("-laptop").is_err()); + assert!(validate_client_name("laptop_1").is_err()); + assert!(validate_client_name("").is_err()); + } +} diff --git a/easytier-core/src/instance/build_capabilities.rs b/easytier-core/src/instance/build_capabilities.rs index ea331d77..dd7ee736 100644 --- a/easytier-core/src/instance/build_capabilities.rs +++ b/easytier-core/src/instance/build_capabilities.rs @@ -60,16 +60,16 @@ fn validate_snapshot( || runtime.public_ipv6_provider.configured_prefix.is_some(), "public IPv6 services", )?; - require( - VPN_PORTAL_AVAILABLE, - peer.vpn_portal_cidr.is_some(), - "the VPN portal", - )?; Ok(()) } pub(super) fn validate(config: &CoreInstanceConfig) -> anyhow::Result<()> { - validate_snapshot(&config.connectivity.runtime, &config.peer.snapshot) + validate_snapshot(&config.connectivity.runtime, &config.peer.snapshot)?; + require( + VPN_PORTAL_AVAILABLE, + config.vpn_portal.is_some(), + "the VPN portal", + ) } pub(super) fn validate_runtime( diff --git a/easytier-core/src/instance/config.rs b/easytier-core/src/instance/config.rs index ab6c38b3..4d762198 100644 --- a/easytier-core/src/instance/config.rs +++ b/easytier-core/src/instance/config.rs @@ -15,6 +15,7 @@ use crate::{ manual::{ManualConnectorOptions, discovery::ManualEndpointDiscoveryConfig}, stun::StunServerConfig, }, + gateway::vpn_portal::{PortalClientConfig, PortalRuntimeConfig}, listener::plan::ListenerRuntimeConfig, packet::CompressorAlgo, peers::{ @@ -231,10 +232,6 @@ impl CoreInstanceConfig { host_routing: host.host_routing, acl: acl.clone(), easytier_version: host.easytier_version.clone(), - vpn_portal_cidr: (!host.ignore_unsupported_config || host.vpn_portal_enabled) - .then(|| config.get_vpn_portal_config()) - .flatten() - .map(|portal| portal.client_cidr), pinned_peers: peers .iter() .cloned() @@ -328,6 +325,20 @@ impl CoreInstanceConfig { Ok(Self { instance_name: config.get_inst_name(), peer, + vpn_portal: (!host.ignore_unsupported_config || host.vpn_portal_enabled) + .then(|| config.get_vpn_portal_config()) + .flatten() + .map(|config| PortalRuntimeConfig { + clients: config + .clients + .into_iter() + .map(|client| PortalClientConfig { + name: client.name, + virtual_ip: client.virtual_ip, + groups: client.groups, + }) + .collect(), + }), connectivity: CoreConnectivityConfig { initial_peers: peers.into_iter().map(|peer| peer.uri).collect(), listeners, diff --git a/easytier-core/src/instance/mod.rs b/easytier-core/src/instance/mod.rs index 5036f46a..b7b6b60c 100644 --- a/easytier-core/src/instance/mod.rs +++ b/easytier-core/src/instance/mod.rs @@ -35,7 +35,6 @@ use url::Url; #[cfg(feature = "tcp-hole-punch")] use crate::connectivity::hole_punch::tcp::TcpHolePunchConnector; use crate::{ - config::peers::{AclRuleConfig, PeerRuntimeSnapshot}, config::runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfig, CoreRuntimeConfigStore}, config::toml::TomlConfig, connectivity::hole_punch::port_mapping::UdpPortMappingPlatform, @@ -91,7 +90,8 @@ use crate::gateway::proxy::icmp_host::IcmpProxyHost; #[cfg(feature = "wrapped-transport")] use crate::gateway::proxy::wrapped_transport::WrappedTransportEngines; #[cfg(feature = "vpn-portal")] -use crate::gateway::vpn_portal::VpnPortalHost; +use crate::gateway::vpn_portal::PortalHost; +use crate::gateway::vpn_portal::PortalRuntimeConfig; #[cfg(feature = "public-ipv6-provider")] use crate::peers::public_ipv6::provider::PublicIpv6ProviderPlatform; @@ -105,7 +105,7 @@ use crate::gateway::proxy::service::CoreProxyModule; #[cfg(feature = "wrapped-transport")] use crate::gateway::proxy::wrapped_transport::WrappedTransportProxyModule; #[cfg(feature = "vpn-portal")] -use crate::gateway::vpn_portal::VpnPortalModule; +use crate::gateway::vpn_portal::PortalModule; #[cfg(feature = "proxy-smoltcp-stack")] use crate::gateway::{ DataPlaneRuntime, DataPlaneSession, PortForwardAdapter, Socks5GatewayAdapter, @@ -182,6 +182,8 @@ pub struct CoreInstanceConfig { pub instance_name: String, pub peer: PortablePeerManagerConfig, pub connectivity: CoreConnectivityConfig, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub vpn_portal: Option, } #[cfg(any(test, feature = "test-utils"))] @@ -192,45 +194,10 @@ pub struct PeerRelaySessionSnapshot { pub has_session: bool, } -fn validate_core_instance_config( - config: &CoreInstanceConfig, -) -> anyhow::Result> { - let acl = config.connectivity.runtime.acl.build()?; - build_capabilities::validate(config)?; - Ok(acl) -} - fn proxy_cidr_snapshot(config: &CoreInstanceRuntimeConfig) -> ProxyCidrSnapshot { ProxyCidrSnapshot::from_proxy_networks(&config.peer.runtime.core.routes.proxy_networks) } -fn retain_core_peer_identity( - peer: &mut Arc, - peer_id: crate::config::PeerId, - instance_id: Option<[u8; 16]>, -) { - let peer = Arc::make_mut(peer); - peer.runtime.core.node.peer_id = Some(peer_id); - peer.runtime.core.node.instance_id = instance_id; -} - -fn retain_runtime_owned_peer_state( - current: &CoreInstanceRuntimeConfig, - next: &mut CoreInstanceRuntimeConfig, - peer_id: crate::config::PeerId, -) { - retain_core_peer_identity( - &mut next.peer, - peer_id, - current.peer.runtime.core.node.instance_id, - ); - let next_peer = Arc::make_mut(&mut next.peer); - next_peer.runtime.stun_info = current.peer.runtime.stun_info.clone(); - if current.services.dhcp_ipv4 && next.services.dhcp_ipv4 { - next_peer.runtime.core.routes.ipv4 = current.peer.runtime.core.routes.ipv4.clone(); - } -} - /// Host-owned resources that must be prepared for the complete Instance /// lifetime, such as a native packet interface. #[async_trait::async_trait] @@ -323,7 +290,7 @@ where #[cfg(feature = "public-ipv6-provider")] pub public_ipv6_provider: Option>, #[cfg(feature = "vpn-portal")] - pub vpn_portal: Option>, + pub vpn_portal: Option>, } impl CoreHostAdapters @@ -441,7 +408,7 @@ where #[cfg(feature = "public-ipv6-provider")] public_ipv6_provider: PublicIpv6ProviderRuntime, #[cfg(feature = "vpn-portal")] - vpn_portal: Arc, + vpn_portal: Arc, #[cfg(feature = "proxy-smoltcp-stack")] pub(super) startup_plan: CoreInstanceStartupPlan, pub(super) runtime_config: CoreRuntimeConfigStore, @@ -502,8 +469,10 @@ where host_config: CoreInstanceHostConfig, mut adapters: CoreHostAdapters, ) -> anyhow::Result> { - let initial_acl = validate_core_instance_config(&config)?; + build_capabilities::validate(&config)?; let instance_name = config.instance_name; + #[cfg(feature = "vpn-portal")] + let vpn_portal_config = config.vpn_portal.clone(); let (packet_tx, packet_rx) = host_packet_channel(); let runtime_config = CoreRuntimeConfigStore::new( config.connectivity.runtime.clone(), @@ -542,7 +511,6 @@ where adapters.credential_storage.take(), foreign_rpc_registrar, )?); - peer_manager.reload_acl(initial_acl.as_ref()); let config = config.connectivity; let listener_plan = prepare_listener_plan( config.listeners.as_ref(), @@ -767,12 +735,13 @@ where public_ipv6_runtime, ); #[cfg(feature = "vpn-portal")] - let vpn_portal = VpnPortalModule::new( + let vpn_portal = PortalModule::new( peer_manager.clone(), runtime_config.clone(), + vpn_portal_config, vpn_portal, events.clone(), - ); + )?; #[cfg(feature = "proxy-cidr-monitor")] let proxy_cidr_monitor = ProxyCidrMonitorRuntime::new(proxy_cidr_monitor_enabled, events.clone()); @@ -843,19 +812,6 @@ where self.state.store(state as u8, Ordering::Release); } - async fn reload_acl_config_inner(&self, config: &AclRuleConfig) -> anyhow::Result<()> { - let acl = config.build()?; - self.peer_manager.reload_acl(acl.as_ref()); - #[cfg(feature = "test-utils")] - self.acl_reload_count.fetch_add(1, Ordering::Relaxed); - Ok(()) - } - - fn sync_peer_runtime_state(&self, snapshot: &PeerRuntimeSnapshot) { - self.peer_manager - .set_avoid_relay_data_preference(snapshot.avoid_relay_data_preference); - } - /// Publishes one complete instance configuration version. Host changes have /// no effect until submitted through this method. pub async fn update_runtime_config( @@ -877,27 +833,15 @@ where anyhow::bail!("runtime config cannot update while instance is stopping or stopped"); } self.validate_runtime_config_capabilities(&config)?; - let current = self.runtime_config.snapshot(); - let refresh_acl_groups = current.peer.peer_group_memberships - != config.peer.peer_group_memberships - || current.peer.acl_group_declarations != config.peer.acl_group_declarations; - if current.services.acl != config.services.acl { - self.reload_acl_config_inner(&config.services.acl).await?; + #[cfg(feature = "test-utils")] + let reload_acl = self.runtime_config.snapshot().services.acl != config.services.acl; + let published = self.peer_manager.update_runtime_config(config).await?; + #[cfg(feature = "test-utils")] + if reload_acl { + self.acl_reload_count.fetch_add(1, Ordering::Relaxed); } - // Foreign-network watchers read this state after the runtime-config - // notification, so publish it before replacing the watched snapshot. - self.sync_peer_runtime_state(&config.peer); - let peer_id = self.peer_id(); - let published = self - .runtime_config - .replace_with_current(config, |current, next| { - retain_runtime_owned_peer_state(current, next, peer_id); - }); self.proxy_cidr_table .update_snapshot(proxy_cidr_snapshot(&published)); - if refresh_acl_groups { - self.refresh_acl_groups().await; - } #[cfg(feature = "proxy-smoltcp-stack")] self.port_forward_adapter .reload( @@ -916,7 +860,10 @@ where &self, config: &CoreInstanceRuntimeConfig, ) -> anyhow::Result<()> { - build_capabilities::validate_runtime(config) + build_capabilities::validate_runtime(config)?; + #[cfg(feature = "vpn-portal")] + self.vpn_portal.validate_runtime_config(config)?; + Ok(()) } pub async fn wait(&self) { diff --git a/easytier-core/src/instance/tests.rs b/easytier-core/src/instance/tests.rs index 569963e1..881a71c4 100644 --- a/easytier-core/src/instance/tests.rs +++ b/easytier-core/src/instance/tests.rs @@ -1,5 +1,3 @@ -use std::sync::Arc; - use async_trait::async_trait; use super::*; @@ -59,21 +57,6 @@ impl ExternalListenerFactory<()> for TestExternalListenerFactory { } } -#[test] -fn runtime_updates_retain_core_owned_peer_identity() { - let mut snapshot = Arc::new(PeerRuntimeSnapshot::default()); - Arc::make_mut(&mut snapshot).runtime.core.node.peer_id = Some(17); - Arc::make_mut(&mut snapshot).runtime.core.node.instance_id = Some([1; 16]); - let submitted = snapshot.clone(); - - retain_core_peer_identity(&mut snapshot, 23, Some([2; 16])); - - assert_eq!(snapshot.runtime.core.node.peer_id, Some(23)); - assert_eq!(snapshot.runtime.core.node.instance_id, Some([2; 16])); - assert_eq!(submitted.runtime.core.node.peer_id, Some(17)); - assert_eq!(submitted.runtime.core.node.instance_id, Some([1; 16])); -} - #[test] fn core_plans_transport_and_external_listener_capabilities() { let self_id = uuid::Uuid::new_v4(); @@ -190,6 +173,7 @@ fn core_instance_config_round_trips_as_normalized_json() { instance_name: String::new(), peer, connectivity: CoreConnectivityConfig::default(), + vpn_portal: None, }; let mut config = config; @@ -228,29 +212,13 @@ fn wasi_create_config_uses_shared_toml() { } #[test] -fn core_instance_config_validation_rejects_invalid_acl_whitelist() { - let peer = crate::peers::peer_manager::PortablePeerManagerConfig::new( - crate::config::peers::PeerRuntimeConfig { - core: crate::config::CoreConfig::default(), - network_identity: crate::config::NetworkIdentity { - network_name: "default".to_owned(), - network_secret: Some("test".to_owned()), - network_secret_digest: None, - }, - stun_info: crate::proto::common::StunInfo::default(), - feature_flags: crate::proto::common::PeerFeatureFlag::default(), - secure_mode: None, - host_routing: crate::config::peers::HostRoutingPolicy::default(), - }, - ); - let mut config = CoreInstanceConfig { - instance_name: String::new(), - peer, - connectivity: CoreConnectivityConfig::default(), +fn acl_config_rejects_invalid_whitelist() { + let config = crate::config::peers::AclRuleConfig { + tcp_whitelist: vec!["9000-8000".to_owned()], + ..Default::default() }; - config.connectivity.runtime.acl.tcp_whitelist = vec!["9000-8000".to_owned()]; - let error = validate_core_instance_config(&config).unwrap_err(); + let error = config.build().unwrap_err(); assert!(error.to_string().contains("Start port must be <= end port")); } @@ -359,8 +327,27 @@ mod portable_runtime { instance_name: network_name.to_owned(), peer, connectivity, + vpn_portal: None, } } + #[cfg(feature = "vpn-portal")] + fn portal_test_config(network_name: &str) -> CoreInstanceConfig { + let mut config = test_config(network_name); + config.peer.snapshot.runtime.network_identity.network_secret = + Some("portal-network-secret".to_owned()); + config.peer.snapshot.runtime.core.routes.ipv4 = Some(IpPrefix { + address: "10.82.0.1".parse().unwrap(), + prefix_len: 24, + }); + config.vpn_portal = Some(crate::gateway::vpn_portal::PortalRuntimeConfig { + clients: vec![crate::gateway::vpn_portal::PortalClientConfig { + name: "alice".to_owned(), + virtual_ip: "10.82.0.2".parse().unwrap(), + groups: Vec::new(), + }], + }); + config + } fn runtime_snapshot(config: &CoreInstanceConfig) -> CoreInstanceRuntimeConfig { CoreInstanceRuntimeConfig { @@ -465,6 +452,71 @@ mod portable_runtime { fn build_instance(config: CoreInstanceConfig) -> anyhow::Result>> { build_with_engines(config, WrappedTransportEngines::default()) } + #[cfg(feature = "vpn-portal")] + #[tokio::test] + async fn runtime_update_rejects_portal_client_address_conflict() { + let instance = build_instance(portal_test_config("portal-runtime-update")).unwrap(); + let before = instance.runtime_config.snapshot(); + let mut conflicting = before.as_ref().clone(); + Arc::make_mut(&mut conflicting.peer) + .runtime + .core + .routes + .ipv4 = Some(IpPrefix { + address: "10.82.0.2".parse().unwrap(), + prefix_len: 24, + }); + + let error = instance + .update_runtime_config(conflicting) + .await + .unwrap_err(); + + assert!( + error.to_string().contains("unusable virtual IP"), + "unexpected runtime update error: {error:#}" + ); + assert_eq!( + instance + .runtime_config + .snapshot() + .peer + .runtime + .core + .routes + .ipv4, + before.peer.runtime.core.routes.ipv4 + ); + instance.peer_manager.clear_resources().await; + } + #[cfg(feature = "vpn-portal")] + #[tokio::test] + async fn runtime_update_allows_removing_portal_client_acl_group() { + let mut config = portal_test_config("portal-acl-group-removal"); + config.vpn_portal.as_mut().unwrap().clients[0].groups = vec!["ops".to_owned()]; + config.peer.snapshot.acl_group_declarations = + vec![crate::config::peers::PeerGroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }]; + let instance = build_instance(config).unwrap(); + let mut updated = instance.runtime_config.snapshot().as_ref().clone(); + Arc::make_mut(&mut updated.peer) + .acl_group_declarations + .clear(); + + instance.update_runtime_config(updated).await.unwrap(); + + assert!( + instance + .runtime_config + .snapshot() + .peer + .acl_group_declarations + .is_empty() + ); + instance.peer_manager.clear_resources().await; + } #[cfg(feature = "dhcp-ipv4")] #[tokio::test] @@ -745,6 +797,70 @@ hostname = "core-owned-config" ); instance.stop().await; } + #[cfg(all(feature = "management", feature = "vpn-portal"))] + #[tokio::test] + async fn rejected_portal_address_patch_does_not_commit_shared_toml() { + use easytier_proto::api::config::InstanceConfigPatch; + + let (packet_sink, _packet_receiver) = tokio::sync::mpsc::channel(16); + let instance = CoreInstance::from_toml( + crate::config::toml::TomlConfig::new_from_str( + r#" +instance_name = "rejected-portal-address-patch" +ipv4 = "10.82.0.1/24" + +[network_identity] +network_name = "rejected-portal-address-patch" +network_secret = "portal-network-secret" + +[vpn_portal_config] +wireguard_listen = "0.0.0.0:51820" + +[[vpn_portal_config.clients]] +name = "alice" +virtual_ip = "10.82.0.2" +"#, + ) + .unwrap(), + adapters(None, Arc::new(packet_sink)), + ) + .unwrap(); + instance.set_state(CoreInstanceState::Running); + + let error = crate::management::apply_config_patch( + &instance, + InstanceConfigPatch { + ipv4: Some("10.82.0.2/24".parse::().unwrap().into()), + ..Default::default() + }, + ) + .await + .unwrap_err(); + + assert!( + error.to_string().contains("unusable virtual IP"), + "unexpected config patch error: {error:#}" + ); + assert_eq!( + instance.toml_config().unwrap().get_ipv4().unwrap(), + "10.82.0.1/24".parse().unwrap() + ); + assert_eq!( + instance + .runtime_config + .snapshot() + .peer + .runtime + .core + .routes + .ipv4 + .as_ref() + .unwrap() + .address, + "10.82.0.1".parse::().unwrap() + ); + instance.peer_manager.clear_resources().await; + } #[cfg(feature = "web-client")] #[tokio::test] diff --git a/easytier-core/src/instance/vpn_portal_extension.rs b/easytier-core/src/instance/vpn_portal_extension.rs index 10e26eef..a46d5630 100644 --- a/easytier-core/src/instance/vpn_portal_extension.rs +++ b/easytier-core/src/instance/vpn_portal_extension.rs @@ -1,5 +1,5 @@ use crate::{ - gateway::vpn_portal::VpnPortalInfoSnapshot, + gateway::vpn_portal::PortalInfoSnapshot, instance::{CoreInstance, CoreInstanceHost}, }; @@ -7,7 +7,7 @@ impl CoreInstance where H: CoreInstanceHost, { - pub async fn vpn_portal_info(&self) -> VpnPortalInfoSnapshot { + pub async fn vpn_portal_info(&self) -> PortalInfoSnapshot { self.vpn_portal.info_snapshot().await } } diff --git a/easytier-core/src/management/full/instance_info.rs b/easytier-core/src/management/full/instance_info.rs index 18797a35..0c471f22 100644 --- a/easytier-core/src/management/full/instance_info.rs +++ b/easytier-core/src/management/full/instance_info.rs @@ -9,6 +9,7 @@ use crate::{ }; /// Builds the process-level running snapshot directly from one core Instance. +#[allow(deprecated)] pub async fn network_instance_running_info( instance: &CoreInstance, ) -> anyhow::Result @@ -50,10 +51,6 @@ where .map(Into::into) .collect::>(); let peer_route_pairs = list_peer_route_pair(peers.clone(), routes.clone()); - #[cfg(feature = "vpn-portal")] - let vpn_portal_cfg = Some(instance.vpn_portal_info().await.client_config); - #[cfg(not(feature = "vpn-portal"))] - let vpn_portal_cfg = Some(String::new()); let dev_name = instance .toml_config() .map(|config| config.get_flags().dev_name) @@ -68,7 +65,8 @@ where ips: Some(node.ip_list), stun_info: Some(node.stun_info), listeners: node.listeners.into_iter().map(Into::into).collect(), - vpn_portal_cfg, + // Client private keys are returned only by the explicit portal RPC. + vpn_portal_cfg: None, peer_id: node.peer_id, }), events: instance.management_events(), diff --git a/easytier-core/src/management/instance_rpc/full.rs b/easytier-core/src/management/instance_rpc/full.rs index 09cc5658..8d66334f 100644 --- a/easytier-core/src/management/instance_rpc/full.rs +++ b/easytier-core/src/management/instance_rpc/full.rs @@ -13,7 +13,8 @@ use easytier_proto::{ ListPortForwardRequest, ListPortForwardResponse, MappedListener, MappedListenerManageRpc, MetricSnapshot, PeerManageRpc, PortForwardManageRpc, RevokeCredentialRequest, RevokeCredentialResponse, StatsRpc, UpsertCredentialRequest, - UpsertCredentialResponse, VpnPortalInfo, VpnPortalRpc, + UpsertCredentialResponse, VpnPortalClientInfo, VpnPortalClientState, VpnPortalInfo, + VpnPortalRpc, }, }, common::PortForwardConfigPb, @@ -26,6 +27,7 @@ use easytier_proto::{ use crate::{ config::toml::ConfigLoader as _, + gateway::vpn_portal::{PortalClientState, PortalInfoSnapshot}, instance::{ CoreInstance, CoreInstanceHost, manager::{InstanceFactory, InstanceManager}, @@ -140,6 +142,47 @@ where } } +fn vpn_portal_info_to_proto(info: PortalInfoSnapshot) -> VpnPortalInfo { + let client_config = info + .clients + .first() + .map(|client| client.client_config.clone()) + .unwrap_or_default(); + let connected_clients = info + .clients + .iter() + .filter(|client| client.state == PortalClientState::Online) + .filter_map(|client| client.endpoint.clone()) + .collect(); + #[allow(deprecated)] + VpnPortalInfo { + vpn_type: info.vpn_type, + client_config, + connected_clients, + clients: info + .clients + .into_iter() + .map(|client| VpnPortalClientInfo { + name: client.name, + virtual_ip: client.virtual_ip.to_string(), + groups: client.groups, + state: match client.state { + PortalClientState::Offline => VpnPortalClientState::Offline as i32, + PortalClientState::Connecting => VpnPortalClientState::Connecting as i32, + PortalClientState::Online => VpnPortalClientState::Online as i32, + PortalClientState::Error => VpnPortalClientState::Error as i32, + }, + peer_id: client.peer_id, + endpoint: client.endpoint, + tunnel_ip: client.tunnel_ip.map(|address| address.to_string()), + client_config: client.client_config, + error: client.error, + }) + .collect(), + listener: info.listener, + } +} + #[async_trait::async_trait] impl VpnPortalRpc for InstanceManagementRpc where @@ -158,11 +201,7 @@ where .vpn_portal_info() .await; Ok(GetVpnPortalInfoResponse { - vpn_portal_info: Some(VpnPortalInfo { - vpn_type: info.vpn_type, - client_config: info.client_config, - connected_clients: info.connected_clients, - }), + vpn_portal_info: Some(vpn_portal_info_to_proto(info)), }) } } @@ -387,3 +426,69 @@ where Err(anyhow::anyhow!("not implemented for management API").into()) } } + +#[cfg(test)] +mod tests { + use std::net::Ipv4Addr; + + use super::*; + use crate::gateway::vpn_portal::PortalClientInfoSnapshot; + + #[test] + #[allow(deprecated)] + fn vpn_portal_info_to_proto_preserves_per_client_status() { + let info = vpn_portal_info_to_proto(PortalInfoSnapshot { + vpn_type: "wireguard".to_owned(), + clients: vec![ + PortalClientInfoSnapshot { + name: "alice".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 2), + groups: vec!["ops".to_owned()], + state: PortalClientState::Online, + peer_id: Some(42), + endpoint: Some("198.51.100.2:51820".to_owned()), + tunnel_ip: Some(Ipv4Addr::new(192, 0, 2, 1)), + client_config: "[Interface]\nPrivateKey = secret\n".to_owned(), + error: None, + }, + PortalClientInfoSnapshot { + name: "bob".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 3), + groups: vec!["guests".to_owned()], + state: PortalClientState::Error, + peer_id: None, + endpoint: None, + tunnel_ip: None, + client_config: String::new(), + error: Some("activation failed".to_owned()), + }, + ], + listener: Some("udp://0.0.0.0:51820".to_owned()), + }); + + assert_eq!(info.vpn_type, "wireguard"); + assert_eq!(info.listener.as_deref(), Some("udp://0.0.0.0:51820")); + assert_eq!(info.client_config, "[Interface]\nPrivateKey = secret\n"); + assert_eq!(info.connected_clients, ["198.51.100.2:51820"]); + assert_eq!(info.clients.len(), 2); + + let online = &info.clients[0]; + assert_eq!(online.name, "alice"); + assert_eq!(online.virtual_ip, "10.82.0.2"); + assert_eq!(online.groups, ["ops"]); + assert_eq!(online.state, VpnPortalClientState::Online as i32); + assert_eq!(online.peer_id, Some(42)); + assert_eq!(online.endpoint.as_deref(), Some("198.51.100.2:51820")); + assert_eq!(online.tunnel_ip.as_deref(), Some("192.0.2.1")); + assert_eq!(online.client_config, "[Interface]\nPrivateKey = secret\n"); + assert_eq!(online.error, None); + + let failed = &info.clients[1]; + assert_eq!(failed.name, "bob"); + assert_eq!(failed.state, VpnPortalClientState::Error as i32); + assert_eq!(failed.peer_id, None); + assert_eq!(failed.endpoint, None); + assert_eq!(failed.tunnel_ip, None); + assert_eq!(failed.error.as_deref(), Some("activation failed")); + } +} diff --git a/easytier-core/src/peers/attached.rs b/easytier-core/src/peers/attached.rs new file mode 100644 index 00000000..fea269d1 --- /dev/null +++ b/easytier-core/src/peers/attached.rs @@ -0,0 +1,1255 @@ +//! Locally attached peers backed by their own portable peer runtime. +//! +//! An attached peer participates in routing, ACL identity, and peer lifecycle +//! exactly like a remote EasyTier node. The only privileged part of the seam is +//! the authenticated in-process connection to another peer manager. + +use std::{ + net::{IpAddr, Ipv4Addr}, + sync::Arc, +}; + +use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; +use parking_lot::Mutex as StdMutex; +use tokio::{runtime::Handle, sync::Mutex, task::JoinHandle}; +use tokio_util::sync::CancellationToken; +use x25519_dalek::{PublicKey, StaticSecret}; + +use crate::{ + config::{ + IpPrefix, PeerId, + peers::PeerRuntimeSnapshot, + runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfig, CoreRuntimeConfigStore}, + }, + host::packet::{HostPacket, HostPacketReceiver, host_packet_channel}, + packet::ZCPacket, + peers::{ + conn::peer_conn::PeerConnId, + peer_manager::{PeerManagerCore, PortablePeerManagerConfig, RouteAlgoType}, + public_ipv6::CorePublicIpv6Runtime, + }, + tunnel::ring::create_ring_tunnel_pair, +}; + +/// Configuration for one locally attached peer. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct AttachedPeerConfig { + pub name: String, + pub virtual_ip: Ipv4Addr, + pub groups: Vec, + /// Stable identity key supplied by the caller; changing it changes the + /// peer identity, so callers must persist and reuse it across restarts. + pub identity_private_key: [u8; 32], +} + +#[derive(Debug, Clone, Copy)] +struct AttachedConnections { + network_peer_id: PeerId, + network_conn_id: PeerConnId, + attached_peer_id: PeerId, + attached_conn_id: PeerConnId, +} +#[cfg(test)] +#[derive(Default)] +struct CleanupPause { + started: tokio::sync::Notify, + resume: tokio::sync::Notify, + finished: tokio::sync::Notify, +} + +struct AttachedCredentialRegistration { + network_peer_manager: Arc, + credential_id: Option, + policy_task: Option>, +} + +impl AttachedCredentialRegistration { + fn register( + network_peer_manager: Arc, + network_runtime_config: CoreRuntimeConfigStore, + public_key: [u8; 32], + configured_groups: Vec, + ) -> anyhow::Result { + let mut peer_changes = network_runtime_config.subscribe_peer_runtime_changes(); + let groups = effective_credential_groups( + network_runtime_config.snapshot().as_ref(), + &configured_groups, + ); + let credential_id = network_peer_manager.register_ephemeral_credential( + public_key, + groups, + false, + Vec::new(), + false, + )?; + let task_peer_manager = network_peer_manager.clone(); + let policy_task = tokio::spawn(async move { + while peer_changes.changed().await.is_ok() { + let _ = peer_changes.borrow_and_update(); + let groups = effective_credential_groups( + network_runtime_config.snapshot().as_ref(), + &configured_groups, + ); + if task_peer_manager + .update_ephemeral_credential_groups(credential_id, groups) + .await + .is_none() + { + return; + } + } + }); + Ok(Self { + network_peer_manager, + credential_id: Some(credential_id), + policy_task: Some(policy_task), + }) + } + + async fn close(mut self) { + if let Some(policy_task) = self.policy_task.take() { + policy_task.abort(); + let _ = policy_task.await; + } + if let Some(credential_id) = self.credential_id.take() { + self.network_peer_manager + .revoke_ephemeral_credential_and_refresh(credential_id) + .await; + } + } + + fn revoke(&mut self) { + if let Some(credential_id) = self.credential_id.take() { + self.network_peer_manager + .revoke_ephemeral_credential(credential_id); + } + } +} + +impl Drop for AttachedCredentialRegistration { + fn drop(&mut self) { + if let Some(policy_task) = self.policy_task.take() { + policy_task.abort(); + } + self.revoke(); + } +} +struct AttachedCleanupResources { + connections: AttachedConnections, + credential_registration: Option, +} + +fn effective_credential_groups( + network: &CoreInstanceRuntimeConfig, + configured_groups: &[String], +) -> Vec { + network + .peer + .acl_group_declarations + .iter() + .filter(|declaration| configured_groups.contains(&declaration.group_name)) + .map(|declaration| declaration.group_name.clone()) + .collect() +} + +fn build_attached_services( + network: &CoreRuntimeConfig, + credential_peer: bool, +) -> CoreRuntimeConfig { + let mut services = network.clone(); + if credential_peer { + services.acl = services.acl.for_credential_peer(); + } + services +} + +/// One complete EasyTier peer connected to another manager in process. +/// +/// The runtime owns the attached manager's lifecycle. Callers exchange only +/// raw IPv4 packets through the Host packet seam. +pub struct AttachedPeerRuntime { + network_peer_manager: Arc, + peer_manager: Arc, + packet_receiver: Mutex, + cleanup_resources: StdMutex>, + cleanup_done: CancellationToken, + runtime_handle: Handle, + closed: CancellationToken, + #[cfg(test)] + cleanup_pause: StdMutex>>, +} + +impl AttachedPeerRuntime { + /// Constructs, starts, and transactionally connects one peer manager. + pub async fn connect( + network_peer_manager: Arc, + network_runtime_config: CoreRuntimeConfigStore, + config: AttachedPeerConfig, + ) -> anyhow::Result> { + let runtime_handle = Handle::current(); + let network = network_runtime_config.snapshot(); + let (peer_snapshot, credential_public_key) = build_peer_snapshot(&network, &config)?; + let credential_registration = credential_public_key + .map(|public_key| { + AttachedCredentialRegistration::register( + network_peer_manager.clone(), + network_runtime_config.clone(), + public_key, + config.groups.clone(), + ) + }) + .transpose()?; + let services = build_attached_services(&network.services, credential_public_key.is_some()); + let runtime_config = CoreRuntimeConfigStore::new(services, Arc::new(peer_snapshot.clone())); + let (packet_sender, packet_receiver) = host_packet_channel(); + let public_ipv6_runtime = + CorePublicIpv6Runtime::new(runtime_config.clone(), Arc::new(()), Arc::new(())); + let flags = peer_snapshot.flags.clone(); + let peer_manager = Arc::new(PeerManagerCore::new( + PortablePeerManagerConfig { + snapshot: peer_snapshot, + route_algo: RouteAlgoType::Ospf, + exit_nodes: Vec::new(), + foreign_context_default_flags: flags, + }, + runtime_config, + Arc::new(()), + packet_sender, + public_ipv6_runtime, + Arc::new(()), + None, + Arc::new(()), + )?); + if let Err(error) = peer_manager + .follow_network_policy(network_runtime_config, config.groups) + .await + { + peer_manager.clear_resources().await; + return Err(error); + } + if let Err(error) = peer_manager.run().await { + peer_manager.clear_resources().await; + return Err(error.into()); + } + + let (network_tunnel, attached_tunnel) = create_ring_tunnel_pair(); + let (network_result, attached_result) = tokio::join!( + network_peer_manager.add_attached_ring_tunnel_as_server(network_tunnel), + peer_manager.add_attached_ring_client_tunnel(attached_tunnel), + ); + let connections = match (network_result, attached_result) { + (Ok((attached_peer_id, network_conn_id)), Ok((network_peer_id, attached_conn_id))) + if attached_peer_id == peer_manager.my_peer_id() + && network_peer_id == network_peer_manager.my_peer_id() => + { + AttachedConnections { + network_peer_id, + network_conn_id, + attached_peer_id, + attached_conn_id, + } + } + (network_result, attached_result) => { + let network_connection = network_result.as_ref().ok().copied(); + let attached_connection = attached_result.as_ref().ok().copied(); + cleanup_partial( + network_peer_manager, + peer_manager.clone(), + network_connection, + attached_connection, + ) + .await; + anyhow::bail!( + "failed to connect attached peer: \ + network={network_result:?}, attached={attached_result:?}" + ); + } + }; + + Ok(Arc::new(Self { + network_peer_manager, + peer_manager, + packet_receiver: Mutex::new(packet_receiver), + cleanup_resources: StdMutex::new(Some(AttachedCleanupResources { + connections, + credential_registration, + })), + cleanup_done: CancellationToken::new(), + runtime_handle, + closed: CancellationToken::new(), + #[cfg(test)] + cleanup_pause: StdMutex::new(None), + })) + } + + pub fn peer_id(&self) -> PeerId { + self.peer_manager.my_peer_id() + } + #[cfg(test)] + fn credential_id(&self) -> Option { + self.cleanup_resources + .lock() + .as_ref() + .and_then(|resources| resources.credential_registration.as_ref()) + .and_then(|registration| registration.credential_id) + } + + /// Receives one packet addressed to the attached peer after routing and ACL. + pub async fn recv_packet(&self) -> Option { + let mut packet_receiver = self.packet_receiver.lock().await; + tokio::select! { + biased; + _ = self.closed.cancelled() => None, + packet = packet_receiver.recv() => packet, + } + } + + /// Injects one raw IPv4 packet as traffic originating from the attached peer. + pub async fn send_packet(&self, payload: &[u8]) -> anyhow::Result<()> { + if self.closed.is_cancelled() { + anyhow::bail!("attached peer is closed"); + } + let destination = ipv4_destination(payload)?; + self.peer_manager + .send_msg_by_ip( + ZCPacket::new_with_payload(payload), + IpAddr::V4(destination), + false, + ) + .await?; + Ok(()) + } + + fn start_cleanup(&self) { + self.closed.cancel(); + let Some(resources) = self.cleanup_resources.lock().take() else { + return; + }; + let AttachedCleanupResources { + connections, + credential_registration, + } = resources; + let network_peer_manager = self.network_peer_manager.clone(); + let peer_manager = self.peer_manager.clone(); + let cleanup_done = self.cleanup_done.clone().drop_guard(); + #[cfg(test)] + let cleanup_pause = self.cleanup_pause.lock().clone(); + self.runtime_handle.spawn(async move { + let _cleanup_done = cleanup_done; + let connection_cleanup = async move { + #[cfg(test)] + if let Some(pause) = cleanup_pause.as_ref() { + pause.started.notify_one(); + pause.resume.notified().await; + } + cleanup(network_peer_manager, peer_manager, connections).await; + #[cfg(test)] + if let Some(pause) = cleanup_pause { + pause.finished.notify_one(); + } + }; + if let Some(credential_registration) = credential_registration { + tokio::join!(connection_cleanup, credential_registration.close()); + } else { + connection_cleanup.await; + } + }); + } + + /// Idempotently tears down both connections and the attached manager. + pub async fn close(&self) { + self.start_cleanup(); + self.cleanup_done.cancelled().await; + } +} + +impl Drop for AttachedPeerRuntime { + fn drop(&mut self) { + self.start_cleanup(); + } +} + +fn ipv4_destination(payload: &[u8]) -> anyhow::Result { + if payload.len() < 20 || payload[0] >> 4 != 4 { + anyhow::bail!("invalid attached-peer IPv4 header"); + } + let header_len = usize::from(payload[0] & 0x0f) * 4; + let total_len = usize::from(u16::from_be_bytes([payload[2], payload[3]])); + if header_len < 20 || header_len > payload.len() || total_len != payload.len() { + anyhow::bail!("invalid attached-peer IPv4 length"); + } + Ok(Ipv4Addr::new( + payload[16], + payload[17], + payload[18], + payload[19], + )) +} + +fn build_peer_snapshot( + network: &CoreInstanceRuntimeConfig, + config: &AttachedPeerConfig, +) -> anyhow::Result<(PeerRuntimeSnapshot, Option<[u8; 32]>)> { + network + .peer + .runtime + .network_identity + .network_secret + .as_deref() + .filter(|secret| !secret.is_empty()) + .ok_or_else(|| anyhow::anyhow!("attached peers require a non-empty network secret"))?; + if network.services.dhcp_ipv4 { + anyhow::bail!("attached peers require a static network-manager IPv4 address"); + } + let network_ipv4 = network + .peer + .runtime + .core + .routes + .ipv4 + .as_ref() + .ok_or_else(|| anyhow::anyhow!("attached peers require a network-manager IPv4 prefix"))?; + let IpAddr::V4(network_address) = network_ipv4.address else { + anyhow::bail!("attached peers require a network-manager IPv4 prefix"); + }; + let network_prefix = cidr::Ipv4Inet::new(network_address, network_ipv4.prefix_len) + .map_err(|error| anyhow::anyhow!("invalid network-manager IPv4 prefix: {error}"))?; + let network_prefix = network_prefix.network(); + if config.virtual_ip == network_address + || !network_prefix.contains(&config.virtual_ip) + || config.virtual_ip == network_prefix.first_address() + || config.virtual_ip == network_prefix.last_address() + { + anyhow::bail!("unusable attached-peer IPv4 address: {}", config.virtual_ip); + } + + let mut snapshot = network.peer.as_ref().clone(); + snapshot.runtime.core.node.peer_id = None; + snapshot.runtime.core.node.instance_id = None; + snapshot.runtime.core.node.hostname = Some(config.name.clone()); + snapshot.runtime.core.routes.ipv4 = Some(IpPrefix { + address: IpAddr::V4(config.virtual_ip), + prefix_len: network_ipv4.prefix_len, + }); + snapshot.runtime.core.routes.ipv6 = None; + snapshot.runtime.core.routes.advertised_routes.clear(); + snapshot.runtime.core.routes.proxy_networks.clear(); + snapshot.runtime.core.routes.foreign_networks.clear(); + snapshot.runtime.core.peer_policy.p2p_enabled = false; + snapshot.runtime.core.peer_policy.relay_peer_rpc = false; + snapshot.runtime.core.peer_policy.relay_data = false; + snapshot.runtime.stun_info = Default::default(); + let supports_conn_list_sync = snapshot.runtime.feature_flags.support_conn_list_sync; + snapshot.runtime.feature_flags = Default::default(); + snapshot.runtime.feature_flags.support_conn_list_sync = supports_conn_list_sync; + snapshot.runtime.feature_flags.disable_p2p = true; + snapshot.runtime.feature_flags.need_p2p = false; + snapshot.runtime.feature_flags.avoid_relay_data = true; + snapshot.flags.disable_p2p = true; + snapshot.flags.need_p2p = false; + snapshot.flags.relay_all_peer_rpc = false; + snapshot.flags.disable_relay_data = true; + snapshot.flags.p2p_only = false; + snapshot.pinned_peers.clear(); + snapshot.avoid_relay_data_preference = true; + snapshot.peer_group_memberships.clear(); + + let credential_public_key = if snapshot + .runtime + .secure_mode + .as_ref() + .is_some_and(|secure| secure.enabled) + { + let private = StaticSecret::from(config.identity_private_key); + let public = PublicKey::from(&private); + snapshot.runtime.network_identity.network_secret = None; + snapshot.runtime.network_identity.network_secret_digest = None; + snapshot.hmac_secret_digest = false; + snapshot.acl_group_declarations.clear(); + snapshot.runtime.secure_mode = Some(crate::proto::common::SecureModeConfig { + enabled: true, + local_private_key: Some(BASE64_STANDARD.encode(private.as_bytes())), + local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), + }); + Some(*public.as_bytes()) + } else { + None + }; + Ok((snapshot, credential_public_key)) +} + +async fn cleanup( + network_peer_manager: Arc, + attached_peer_manager: Arc, + connections: AttachedConnections, +) { + cleanup_partial( + network_peer_manager, + attached_peer_manager, + Some((connections.attached_peer_id, connections.network_conn_id)), + Some((connections.network_peer_id, connections.attached_conn_id)), + ) + .await; +} + +async fn cleanup_partial( + network_peer_manager: Arc, + attached_peer_manager: Arc, + network_connection: Option<(PeerId, PeerConnId)>, + attached_connection: Option<(PeerId, PeerConnId)>, +) { + if let (Some((attached_peer_id, network_conn_id)), Some((network_peer_id, attached_conn_id))) = + (network_connection, attached_connection) + { + let _ = tokio::join!( + network_peer_manager.close_peer_conn(attached_peer_id, &network_conn_id), + attached_peer_manager.close_peer_conn(network_peer_id, &attached_conn_id), + ); + } else { + if let Some((attached_peer_id, network_conn_id)) = network_connection { + let _ = network_peer_manager + .close_peer_conn(attached_peer_id, &network_conn_id) + .await; + } + if let Some((network_peer_id, attached_conn_id)) = attached_connection { + let _ = attached_peer_manager + .close_peer_conn(network_peer_id, &attached_conn_id) + .await; + } + } + if let Some((attached_peer_id, _)) = network_connection { + let _ = network_peer_manager + .get_peer_map() + .close_peer(attached_peer_id) + .await; + } + if let Some((network_peer_id, _)) = attached_connection { + let _ = attached_peer_manager + .get_peer_map() + .close_peer(network_peer_id) + .await; + } + attached_peer_manager.clear_resources().await; +} + +#[cfg(test)] +mod tests { + use std::{collections::BTreeSet, time::Duration}; + + use super::*; + use crate::{ + config::{ + CoreConfig, NetworkIdentity, NodeConfig, + peers::{HostRoutingPolicy, PeerGroupIdentity, PeerRuntimeConfig}, + runtime::CoreRuntimeConfig, + }, + proto::{ + acl::{Acl, AclV1, GroupInfo}, + common::{PeerFeatureFlag, StunInfo}, + }, + }; + + fn peer_manager_with_acl( + tcp_whitelist: Vec, + ) -> (Arc, CoreRuntimeConfigStore) { + peer_manager_with_acl_and_secure(tcp_whitelist, false) + } + + fn peer_manager_with_acl_and_secure( + tcp_whitelist: Vec, + secure_mode_enabled: bool, + ) -> (Arc, CoreRuntimeConfigStore) { + let secure_mode = secure_mode_enabled.then(|| { + let private = StaticSecret::from([99u8; 32]); + let public = PublicKey::from(&private); + crate::proto::common::SecureModeConfig { + enabled: true, + local_private_key: Some(BASE64_STANDARD.encode(private.as_bytes())), + local_public_key: Some(BASE64_STANDARD.encode(public.as_bytes())), + } + }); + let runtime = PeerRuntimeConfig { + core: CoreConfig { + node: NodeConfig { + network_name: "attached-test".to_owned(), + ..Default::default() + }, + routes: crate::config::RouteConfig { + ipv4: Some(IpPrefix { + address: IpAddr::V4(Ipv4Addr::new(10, 82, 0, 1)), + prefix_len: 24, + }), + ..Default::default() + }, + ..Default::default() + }, + network_identity: NetworkIdentity::new( + "attached-test".to_owned(), + "shared-secret".to_owned(), + ), + stun_info: StunInfo::default(), + feature_flags: PeerFeatureFlag::default(), + secure_mode, + host_routing: HostRoutingPolicy::default(), + }; + let mut portable = PortablePeerManagerConfig::new(runtime); + portable.snapshot.acl_group_declarations = vec![PeerGroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }]; + let services = CoreRuntimeConfig { + acl: crate::config::peers::AclRuleConfig { + acl: Some(Acl { + acl_v1: Some(AclV1 { + chains: Vec::new(), + group: Some(GroupInfo { + declares: vec![crate::proto::acl::GroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }], + members: Vec::new(), + }), + }), + }), + tcp_whitelist, + ..Default::default() + }, + ..Default::default() + }; + let store = CoreRuntimeConfigStore::new(services, Arc::new(portable.snapshot.clone())); + let (packet_sender, _packet_receiver) = host_packet_channel(); + let public_ipv6_runtime = + CorePublicIpv6Runtime::new(store.clone(), Arc::new(()), Arc::new(())); + let peer_manager = Arc::new( + PeerManagerCore::new( + portable, + store.clone(), + Arc::new(()), + packet_sender, + public_ipv6_runtime, + Arc::new(()), + None, + Arc::new(()), + ) + .unwrap(), + ); + (peer_manager, store) + } + + fn peer_manager() -> (Arc, CoreRuntimeConfigStore) { + peer_manager_with_acl(Vec::new()) + } + + async fn wait_for_route( + peer: &PeerManagerCore, + attached_peer_id: PeerId, + address: Ipv4Addr, + present: bool, + ) { + tokio::time::timeout(Duration::from_secs(10), async { + loop { + let found = peer.list_route_snapshots().await.iter().any(|route| { + route.peer_id == attached_peer_id + && route.ipv4_addr == Some(cidr::Ipv4Inet::new(address, 24).unwrap().into()) + }); + if found == present { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("attached route state did not converge"); + } + + async fn wait_for_acl_groups( + network_peer_manager: &PeerManagerCore, + peer_id: PeerId, + expected: &[&str], + ) { + let expected = expected + .iter() + .map(|group| (*group).to_owned()) + .collect::>(); + tokio::time::timeout(Duration::from_secs(10), async { + loop { + let advertised = network_peer_manager + .get_route() + .get_peer_groups(peer_id) + .iter() + .cloned() + .collect::>(); + if advertised == expected { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("attached ACL groups did not converge"); + } + + async fn wait_for_acl_rules(attached: &AttachedPeerRuntime) { + tokio::time::timeout(Duration::from_secs(10), async { + loop { + if !attached + .peer_manager + .acl_filter() + .get_processor() + .get_rules_stats() + .is_empty() + { + return; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("attached peer manager did not apply updated ACL rules"); + } + + #[tokio::test] + async fn secure_attached_peer_updates_groups_with_declarations() { + let (network_peer_manager, store) = peer_manager_with_acl_and_secure(Vec::new(), true); + network_peer_manager.run().await.unwrap(); + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store.clone(), + AttachedPeerConfig { + name: "group-update".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 2), + groups: vec!["ops".to_owned(), "audit".to_owned()], + identity_private_key: [2; 32], + }, + ) + .await + .unwrap(); + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &["ops"]).await; + store.update_peer_with(|peer| { + peer.acl_group_declarations.push(PeerGroupIdentity { + group_name: "audit".to_owned(), + group_secret: "audit-secret".to_owned(), + }); + }); + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &["audit", "ops"]).await; + + store.update_peer_with(|peer| peer.acl_group_declarations.clear()); + + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &[]).await; + attached.close().await; + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn ephemeral_credential_group_update_refreshes_route() { + let (network_peer_manager, store) = peer_manager_with_acl_and_secure(Vec::new(), true); + network_peer_manager.run().await.unwrap(); + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "direct-group-update".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 2), + groups: vec!["ops".to_owned()], + identity_private_key: [5; 32], + }, + ) + .await + .unwrap(); + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &["ops"]).await; + let credential_id = attached.credential_id().unwrap(); + + assert_eq!( + network_peer_manager + .update_ephemeral_credential_groups(credential_id, Vec::new()) + .await, + Some(true) + ); + + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &[]).await; + attached.close().await; + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn peer_manager_loads_initial_acl_from_its_runtime_config() { + let (peer_manager, _store) = peer_manager_with_acl(vec!["22".to_owned()]); + + assert!( + !peer_manager + .acl_filter() + .get_processor() + .get_rules_stats() + .is_empty() + ); + } + + #[tokio::test] + async fn invalid_acl_update_does_not_publish_manager_config() { + let (peer_manager, store) = peer_manager(); + let before = store.snapshot(); + let mut invalid = before.as_ref().clone(); + invalid.services.acl.tcp_whitelist = vec!["not-a-port".to_owned()]; + + assert!(peer_manager.update_runtime_config(invalid).await.is_err()); + assert_eq!(store.snapshot().services.acl, before.services.acl); + assert!( + peer_manager + .acl_filter() + .get_processor() + .get_rules_stats() + .is_empty() + ); + } + + #[tokio::test] + async fn dropped_credential_registration_revokes_credential() { + let (network_peer_manager, store) = peer_manager_with_acl_and_secure(Vec::new(), true); + let credential_public_key = *PublicKey::from(&StaticSecret::from([4; 32])).as_bytes(); + + let registration = AttachedCredentialRegistration::register( + network_peer_manager.clone(), + store, + credential_public_key, + vec!["ops".to_owned()], + ) + .unwrap(); + assert!( + network_peer_manager + .credential_manager() + .is_pubkey_trusted(&credential_public_key) + ); + + drop(registration); + + assert!( + !network_peer_manager + .credential_manager() + .is_pubkey_trusted(&credential_public_key) + ); + } + + #[tokio::test] + async fn secure_attached_runtime_removes_admin_and_group_secrets() { + let (_network_peer_manager, store) = peer_manager_with_acl_and_secure(Vec::new(), true); + let network = store.snapshot(); + let config = AttachedPeerConfig { + name: "sanitized".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 2), + groups: vec!["ops".to_owned()], + identity_private_key: [3; 32], + }; + + let (snapshot, credential_public_key) = + build_peer_snapshot(network.as_ref(), &config).unwrap(); + let services = build_attached_services(&network.services, credential_public_key.is_some()); + + assert!(credential_public_key.is_some()); + assert!(snapshot.runtime.network_identity.network_secret.is_none()); + assert!( + snapshot + .runtime + .network_identity + .network_secret_digest + .is_none() + ); + assert!(snapshot.peer_group_memberships.is_empty()); + assert!(snapshot.acl_group_declarations.is_empty()); + assert!( + services + .acl + .acl + .as_ref() + .unwrap() + .acl_v1 + .as_ref() + .unwrap() + .group + .is_none() + ); + assert!( + network + .services + .acl + .acl + .as_ref() + .unwrap() + .acl_v1 + .as_ref() + .unwrap() + .group + .is_some() + ); + } + + #[tokio::test] + async fn secure_attached_peer_uses_credential_identity_and_granted_groups() { + let (network_peer_manager, store) = peer_manager_with_acl_and_secure(Vec::new(), true); + network_peer_manager.run().await.unwrap(); + let identity_private_key = [1; 32]; + let credential_public_key = + *PublicKey::from(&StaticSecret::from(identity_private_key)).as_bytes(); + assert!( + !network_peer_manager + .credential_manager() + .is_pubkey_trusted(&credential_public_key) + ); + + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "credential".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 2), + groups: vec!["ops".to_owned()], + identity_private_key, + }, + ) + .await + .unwrap(); + assert!( + network_peer_manager + .credential_manager() + .is_pubkey_trusted(&credential_public_key) + ); + + assert!( + !attached.peer_manager.can_manage_credentials(), + "Secure-Mode attached peer retained administrator credentials" + ); + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &["ops"]).await; + + attached.close().await; + assert!( + !network_peer_manager + .credential_manager() + .is_pubkey_trusted(&credential_public_key) + ); + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn peer_managers_own_acl_updates_for_attached_clients() { + let (network_peer_manager, store) = peer_manager(); + let peer_subscribers_before_run = store.peer_change_subscriber_count(); + network_peer_manager.run().await.unwrap(); + tokio::time::timeout(Duration::from_secs(10), async { + while store.peer_change_subscriber_count() == peer_subscribers_before_run { + tokio::task::yield_now().await; + } + }) + .await + .expect("network peer manager did not subscribe to config updates"); + let peer_subscribers = store.peer_change_subscriber_count(); + let service_subscribers = store.service_change_subscriber_count(); + + let first = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store.clone(), + AttachedPeerConfig { + name: "first".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 2), + groups: vec!["ops".to_owned()], + identity_private_key: [1; 32], + }, + ) + .await + .unwrap(); + let second = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store.clone(), + AttachedPeerConfig { + name: "second".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 3), + groups: Vec::new(), + identity_private_key: [2; 32], + }, + ) + .await + .unwrap(); + + assert_ne!(first.peer_id(), second.peer_id()); + assert!(!Arc::ptr_eq( + &first.peer_manager.acl_filter(), + &second.peer_manager.acl_filter() + )); + assert_eq!(store.peer_change_subscriber_count(), peer_subscribers + 2); + assert_eq!( + store.service_change_subscriber_count(), + service_subscribers + 2 + ); + wait_for_route( + &network_peer_manager, + first.peer_id(), + Ipv4Addr::new(10, 82, 0, 2), + true, + ) + .await; + wait_for_route( + &network_peer_manager, + second.peer_id(), + Ipv4Addr::new(10, 82, 0, 3), + true, + ) + .await; + wait_for_acl_groups(&network_peer_manager, first.peer_id(), &["ops"]).await; + + store.update_peer_with(|peer| peer.acl_group_declarations.clear()); + wait_for_acl_groups(&network_peer_manager, first.peer_id(), &[]).await; + + let mut updated = store.snapshot().as_ref().clone(); + updated.services.acl.tcp_whitelist.push("22".to_owned()); + store.replace(updated); + wait_for_acl_rules(&first).await; + wait_for_acl_rules(&second).await; + + let first_peer_id = first.peer_id(); + first.close().await; + assert_eq!(store.peer_change_subscriber_count(), peer_subscribers + 1); + assert_eq!( + store.service_change_subscriber_count(), + service_subscribers + 1 + ); + wait_for_route( + &network_peer_manager, + first_peer_id, + Ipv4Addr::new(10, 82, 0, 2), + false, + ) + .await; + second.close().await; + assert_eq!(store.peer_change_subscriber_count(), peer_subscribers); + assert_eq!(store.service_change_subscriber_count(), service_subscribers); + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn attached_peer_reconnects_after_acl_group_removal() { + let (network_peer_manager, store) = peer_manager(); + network_peer_manager.run().await.unwrap(); + store.update_peer_with(|peer| peer.acl_group_declarations.clear()); + + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "reconnected".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 4), + groups: vec!["ops".to_owned()], + identity_private_key: [3; 32], + }, + ) + .await + .unwrap(); + + wait_for_route( + &network_peer_manager, + attached.peer_id(), + Ipv4Addr::new(10, 82, 0, 4), + true, + ) + .await; + wait_for_acl_groups(&network_peer_manager, attached.peer_id(), &[]).await; + attached.close().await; + network_peer_manager.clear_resources().await; + } + #[tokio::test] + async fn cancelled_close_can_be_awaited_again() { + let (network_peer_manager, store) = peer_manager(); + network_peer_manager.run().await.unwrap(); + let identity_private_key = [6; 32]; + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "cancelled-close".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 4), + groups: Vec::new(), + identity_private_key, + }, + ) + .await + .unwrap(); + let peer_id = attached.peer_id(); + wait_for_route( + &network_peer_manager, + peer_id, + Ipv4Addr::new(10, 82, 0, 4), + true, + ) + .await; + assert!( + network_peer_manager + .get_peer_map() + .has_direct_attached_peer(peer_id), + "attached connection was not live before close" + ); + + let pause = Arc::new(CleanupPause::default()); + *attached.cleanup_pause.lock() = Some(pause.clone()); + let first_close = tokio::spawn({ + let attached = attached.clone(); + async move { attached.close().await } + }); + pause.started.notified().await; + first_close.abort(); + let _ = first_close.await; + pause.resume.notify_one(); + assert!( + network_peer_manager + .get_peer_map() + .has_direct_attached_peer(peer_id), + "cancelling close unexpectedly completed connection cleanup" + ); + + tokio::time::timeout(Duration::from_secs(5), attached.close()) + .await + .expect("retrying close did not await the in-flight cleanup"); + assert!( + !network_peer_manager + .get_peer_map() + .has_direct_attached_peer(peer_id), + "cancelled close left the attached connection live" + ); + wait_for_route( + &network_peer_manager, + peer_id, + Ipv4Addr::new(10, 82, 0, 4), + false, + ) + .await; + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn closing_attached_peer_unblocks_packet_receiver() { + let (network_peer_manager, store) = peer_manager(); + network_peer_manager.run().await.unwrap(); + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "closed-receiver".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 4), + groups: Vec::new(), + identity_private_key: [3; 32], + }, + ) + .await + .unwrap(); + let waiting = tokio::spawn({ + let attached = attached.clone(); + async move { attached.recv_packet().await } + }); + tokio::task::yield_now().await; + + attached.close().await; + + let packet = tokio::time::timeout(Duration::from_secs(1), waiting) + .await + .expect("packet receiver remained blocked after close") + .unwrap(); + assert!(packet.is_none()); + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn dropping_attached_peer_removes_its_route() { + let (network_peer_manager, store) = peer_manager(); + network_peer_manager.run().await.unwrap(); + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "dropped".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 4), + groups: Vec::new(), + identity_private_key: [3; 32], + }, + ) + .await + .unwrap(); + let peer_id = attached.peer_id(); + wait_for_route( + &network_peer_manager, + peer_id, + Ipv4Addr::new(10, 82, 0, 4), + true, + ) + .await; + + drop(attached); + + wait_for_route( + &network_peer_manager, + peer_id, + Ipv4Addr::new(10, 82, 0, 4), + false, + ) + .await; + network_peer_manager.clear_resources().await; + } + + #[tokio::test] + async fn dropping_secure_attached_peer_refreshes_route_trust_before_disconnect() { + let (network_peer_manager, store) = peer_manager_with_acl_and_secure(Vec::new(), true); + network_peer_manager.run().await.unwrap(); + let identity_private_key = [7; 32]; + let credential_public_key = + *PublicKey::from(&StaticSecret::from(identity_private_key)).as_bytes(); + let attached = AttachedPeerRuntime::connect( + network_peer_manager.clone(), + store, + AttachedPeerConfig { + name: "dropped-secure".to_owned(), + virtual_ip: Ipv4Addr::new(10, 82, 0, 5), + groups: vec!["ops".to_owned()], + identity_private_key, + }, + ) + .await + .unwrap(); + let peer_id = attached.peer_id(); + wait_for_route( + &network_peer_manager, + peer_id, + Ipv4Addr::new(10, 82, 0, 5), + true, + ) + .await; + wait_for_acl_groups(&network_peer_manager, peer_id, &["ops"]).await; + + let pause = Arc::new(CleanupPause::default()); + *attached.cleanup_pause.lock() = Some(pause.clone()); + drop(attached); + pause.started.notified().await; + + tokio::time::timeout(Duration::from_secs(5), async { + while network_peer_manager + .credential_manager() + .is_pubkey_trusted(&credential_public_key) + { + tokio::task::yield_now().await; + } + }) + .await + .expect("dropped attached credential remained trusted"); + wait_for_route( + &network_peer_manager, + peer_id, + Ipv4Addr::new(10, 82, 0, 5), + false, + ) + .await; + assert!( + !network_peer_manager + .get_peer_map() + .has_direct_attached_peer(peer_id), + "credential revocation did not refresh route trust before cleanup resumed" + ); + + pause.resume.notify_one(); + pause.finished.notified().await; + network_peer_manager.clear_resources().await; + } +} diff --git a/easytier-core/src/peers/conn/peer.rs b/easytier-core/src/peers/conn/peer.rs index 04c16aa9..7ab3435d 100644 --- a/easytier-core/src/peers/conn/peer.rs +++ b/easytier-core/src/peers/conn/peer.rs @@ -291,6 +291,13 @@ impl Peer { .any(|entry| !entry.value().is_closed() && !entry.value().is_hole_punched()) } + pub(crate) fn has_direct_attached_conn(&self) -> bool { + self.conns.iter().any(|entry| { + let conn = entry.value(); + !conn.is_closed() && !conn.is_hole_punched() && conn.is_attached() + }) + } + pub fn get_directly_connections(&self) -> DashSet { self.conns .iter() diff --git a/easytier-core/src/peers/conn/peer_conn.rs b/easytier-core/src/peers/conn/peer_conn.rs index c226b4a2..ee829bce 100644 --- a/easytier-core/src/peers/conn/peer_conn.rs +++ b/easytier-core/src/peers/conn/peer_conn.rs @@ -34,9 +34,9 @@ use super::{ peer_session::{PeerSession, PeerSessionAction}, }; use crate::peers::{ - PacketRecvChan, + PacketRecvChan, PeerConnectionOrigin, PeerPacketIngress, context::{ArcPeerContext, NetworkIdentity, NetworkSecretDigest}, - send_packet_to_chan, + send_peer_packet_to_chan, traffic_metrics::data_packet_payload_len, }; use crate::{ @@ -272,6 +272,7 @@ impl PeerConnCloseNotify { pub struct PeerConn { conn_id: PeerConnId, + origin: PeerConnectionOrigin, my_peer_id: PeerId, peer_id_hint: Option, @@ -319,21 +320,30 @@ impl Debug for PeerConn { } impl PeerConn { + #[cfg(test)] pub(crate) fn new( my_peer_id: PeerId, context: ArcPeerContext, tunnel: Box, peer_session_store: Arc, ) -> Self { - Self::new_with_peer_id_hint(my_peer_id, context, tunnel, None, peer_session_store) + Self::new_with_peer_id_hint_and_origin( + my_peer_id, + context, + tunnel, + None, + peer_session_store, + PeerConnectionOrigin::Network, + ) } - pub(crate) fn new_with_peer_id_hint( + pub(crate) fn new_with_peer_id_hint_and_origin( my_peer_id: PeerId, context: ArcPeerContext, tunnel: Box, peer_id_hint: Option, peer_session_store: Arc, + origin: PeerConnectionOrigin, ) -> Self { let flags = context.flags(); let tunnel_info = tunnel.info(); @@ -363,6 +373,7 @@ impl PeerConn { PeerConn { conn_id, + origin, my_peer_id, peer_id_hint, @@ -428,6 +439,10 @@ impl PeerConn { self.conn_id } + pub(crate) fn is_attached(&self) -> bool { + self.origin == PeerConnectionOrigin::Attached + } + pub fn set_is_hole_punched(&mut self, is_hole_punched: bool) { self.is_hole_punched = is_hole_punched; } @@ -1286,6 +1301,11 @@ impl PeerConn { let ctrl_sender = self.ctrl_resp_sender.clone(); let conn_info_for_instrument = self.get_conn_info(); let context = self.context.clone(); + let ingress = PeerPacketIngress::Peer { + peer_id: self.get_peer_id(), + conn_id: self.conn_id, + origin: self.origin, + }; let control_network_name = conn_info_for_instrument.network_name.clone(); let is_foreign_network = @@ -1329,7 +1349,10 @@ impl PeerConn { if let Err(e) = ctrl_sender.send(zc_packet) { tracing::error!(?e, "peer conn send ctrl resp error"); } - } else if send_packet_to_chan(&sender, zc_packet).await.is_err() { + } else if send_peer_packet_to_chan(&sender, zc_packet, ingress) + .await + .is_err() + { break; } diff --git a/easytier-core/src/peers/conn/peer_map.rs b/easytier-core/src/peers/conn/peer_map.rs index 769dc0da..00a6432e 100644 --- a/easytier-core/src/peers/conn/peer_map.rs +++ b/easytier-core/src/peers/conn/peer_map.rs @@ -143,6 +143,11 @@ impl PeerMap { peer_id == self.my_peer_id || self.peer_map.contains_key(&peer_id) } + pub(crate) fn has_direct_attached_peer(&self, peer_id: PeerId) -> bool { + self.get_peer_by_id(peer_id) + .is_some_and(|peer| peer.has_direct_attached_conn()) + } + pub(crate) fn is_self(&self, peer_id: PeerId) -> bool { peer_id == self.my_peer_id } diff --git a/easytier-core/src/peers/context.rs b/easytier-core/src/peers/context.rs index e790facb..a259730e 100644 --- a/easytier-core/src/peers/context.rs +++ b/easytier-core/src/peers/context.rs @@ -76,7 +76,6 @@ pub struct PeerRuntimeSnapshotInput { pub host_routing: HostRoutingPolicy, pub acl: Option, pub easytier_version: String, - pub vpn_portal_cidr: Option, pub pinned_peers: Vec<(url::Url, Option)>, pub ospf_update_my_foreign_network_interval_sec: u64, pub max_direct_conns_per_peer_in_foreign_network: usize, @@ -126,7 +125,6 @@ impl PeerRuntimeSnapshot { host_routing, acl, easytier_version, - vpn_portal_cidr, pinned_peers, ospf_update_my_foreign_network_interval_sec, max_direct_conns_per_peer_in_foreign_network, @@ -183,7 +181,6 @@ impl PeerRuntimeSnapshot { easytier_version, avoid_relay_data_preference, flags, - vpn_portal_cidr, pinned_peers, peer_group_memberships, acl_group_declarations, @@ -331,6 +328,10 @@ impl CorePeerContext { self.config.snapshot().peer.clone() } + pub(crate) fn runtime_config_store(&self) -> CoreRuntimeConfigStore { + self.config.clone() + } + pub fn stats_manager(&self) -> Arc { self.stats_manager.clone() } @@ -604,10 +605,6 @@ pub(crate) trait PeerContext: Send + Sync { Vec::new() } - fn vpn_portal_cidr(&self) -> Option { - None - } - fn hostname(&self) -> String { String::new() } @@ -823,10 +820,6 @@ impl PeerContext for CorePeerContext { self.snapshot().runtime.core.routes.proxy_networks.clone() } - fn vpn_portal_cidr(&self) -> Option { - self.snapshot().vpn_portal_cidr - } - fn hostname(&self) -> String { self.snapshot() .runtime @@ -1131,7 +1124,6 @@ pub(crate) mod tests { }, acl, easytier_version: "host-version".to_owned(), - vpn_portal_cidr: Some("10.30.0.0/24".parse().unwrap()), pinned_peers: vec![( "tcp://192.0.2.10:11010".parse().unwrap(), Some("peer-key".to_owned()), @@ -1211,10 +1203,6 @@ pub(crate) mod tests { assert!(snapshot.avoid_relay_data_preference); assert_eq!(snapshot.easytier_version, "host-version"); - assert_eq!( - snapshot.vpn_portal_cidr, - Some("10.30.0.0/24".parse().unwrap()) - ); assert_eq!( snapshot.pinned_peers, vec![( diff --git a/easytier-core/src/peers/credential_manager.rs b/easytier-core/src/peers/credential_manager.rs index c45b096b..62d87abb 100644 --- a/easytier-core/src/peers/credential_manager.rs +++ b/easytier-core/src/peers/credential_manager.rs @@ -116,6 +116,7 @@ pub trait CredentialStorage: Send + Sync + 'static { pub(crate) struct CredentialManager { credentials: Mutex>, + ephemeral_credentials: Mutex>, storage: Option>, storage_write: Mutex<()>, } @@ -130,6 +131,7 @@ impl CredentialManager { pub fn new() -> Self { Self { credentials: Mutex::new(HashMap::new()), + ephemeral_credentials: Mutex::new(HashMap::new()), storage: None, storage_write: Mutex::new(()), } @@ -149,6 +151,7 @@ impl CredentialManager { }; Self { credentials: Mutex::new(credentials), + ephemeral_credentials: Mutex::new(HashMap::new()), storage: Some(storage), storage_write: Mutex::new(()), } @@ -268,6 +271,72 @@ impl CredentialManager { removed } + pub fn register_ephemeral_credential( + &self, + public_key: [u8; 32], + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + reusable: bool, + ) -> Result { + let entry = CredentialEntry { + pubkey: BASE64_STANDARD.encode(public_key), + secret: String::new(), + groups, + allow_relay, + allowed_proxy_cidrs, + reusable, + expiry_unix: i64::MAX, + created_at_unix: current_unix_timestamp(), + }; + + let _storage_write = self.storage_write.lock().unwrap(); + if self + .credentials + .lock() + .unwrap() + .values() + .any(|existing| existing.pubkey == entry.pubkey) + || self + .ephemeral_credentials + .lock() + .unwrap() + .values() + .any(|existing| existing.pubkey == entry.pubkey) + { + return Err("credential public key is already registered".to_owned()); + } + + let credential_id = uuid::Uuid::new_v4(); + self.ephemeral_credentials + .lock() + .unwrap() + .insert(credential_id, entry); + Ok(credential_id) + } + + pub fn update_ephemeral_credential_groups( + &self, + credential_id: uuid::Uuid, + groups: Vec, + ) -> Option { + let mut credentials = self.ephemeral_credentials.lock().unwrap(); + let credential = credentials.get_mut(&credential_id)?; + if credential.groups == groups { + return Some(false); + } + credential.groups = groups; + Some(true) + } + + pub fn revoke_ephemeral_credential(&self, credential_id: uuid::Uuid) -> bool { + self.ephemeral_credentials + .lock() + .unwrap() + .remove(&credential_id) + .is_some() + } + pub fn upsert_credential(&self, options: CredentialUpsertOptions) -> Result { let CredentialUpsertOptions { credential_id, @@ -310,6 +379,15 @@ impl CredentialManager { }) { return Err("credential_secret is already used by another credential_id".to_string()); } + if self + .ephemeral_credentials + .lock() + .unwrap() + .values() + .any(|existing| existing.pubkey == entry.pubkey) + { + return Err("credential public key is already registered".to_owned()); + } let changed = credentials.get(&credential_id).is_none_or(|existing| { existing.secret != entry.secret || existing.pubkey != entry.pubkey @@ -356,29 +434,43 @@ impl CredentialManager { pub fn get_trusted_pubkeys(&self, network_secret: &str) -> Vec { let now = current_unix_timestamp(); - - self.credentials + let to_proof = |entry: &CredentialEntry| { + entry.to_trusted_credential().map(|credential| { + TrustedCredentialPubkeyProof::new_signed(credential, network_secret) + }) + }; + let mut trusted = self + .credentials .lock() .unwrap() .values() .filter(|entry| entry.is_active_at(now)) - .filter_map(|entry| { - entry.to_trusted_credential().map(|credential| { - TrustedCredentialPubkeyProof::new_signed(credential, network_secret) - }) - }) - .collect() + .filter_map(to_proof) + .collect::>(); + trusted.extend( + self.ephemeral_credentials + .lock() + .unwrap() + .values() + .filter_map(to_proof), + ); + trusted } pub fn is_pubkey_trusted(&self, pubkey: &[u8]) -> bool { let now = current_unix_timestamp(); - let encoded = BASE64_STANDARD.encode(pubkey); self.credentials .lock() .unwrap() .values() .any(|entry| entry.pubkey == encoded && entry.is_active_at(now)) + || self + .ephemeral_credentials + .lock() + .unwrap() + .values() + .any(|entry| entry.pubkey == encoded) } pub fn list_credentials(&self) -> Vec { @@ -725,4 +817,33 @@ mod tests { assert!(manager.list_credentials().is_empty()); } + + #[test] + fn ephemeral_credentials_are_trusted_but_not_persisted_or_listed() { + let storage = Arc::new(MemoryCredentialStorage::default()); + let manager = CredentialManager::from_storage(storage.clone()); + let private = StaticSecret::from([7u8; 32]); + let public = *PublicKey::from(&private).as_bytes(); + + let credential_id = manager + .register_ephemeral_credential(public, vec!["ops".to_owned()], false, Vec::new(), false) + .unwrap(); + + assert!(manager.is_pubkey_trusted(&public)); + let trusted = manager.get_trusted_pubkeys("network-secret"); + assert_eq!(trusted.len(), 1); + let credential = trusted[0].credential.as_ref().unwrap(); + assert_eq!(credential.pubkey, public); + assert_eq!(credential.groups, ["ops"]); + assert!(!credential.allow_relay); + assert!(credential.allowed_proxy_cidrs.is_empty()); + assert_eq!(credential.reusable, Some(false)); + assert!(manager.list_credentials().is_empty()); + assert!(storage.serialized.lock().unwrap().is_none()); + + assert!(manager.revoke_ephemeral_credential(credential_id)); + assert!(!manager.is_pubkey_trusted(&public)); + assert!(manager.get_trusted_pubkeys("network-secret").is_empty()); + assert!(storage.serialized.lock().unwrap().is_none()); + } } diff --git a/easytier-core/src/peers/foreign_network/client.rs b/easytier-core/src/peers/foreign_network/client.rs index cf97dbf7..e5238c92 100644 --- a/easytier-core/src/peers/foreign_network/client.rs +++ b/easytier-core/src/peers/foreign_network/client.rs @@ -93,6 +93,10 @@ impl ForeignNetworkClient { }))); } + pub(crate) fn stop(&self) { + self.task.lock().unwrap().take(); + } + pub fn get_peer_map(&self) -> Arc { self.peer_map.clone() } diff --git a/easytier-core/src/peers/mod.rs b/easytier-core/src/peers/mod.rs index bf36a984..4b796380 100644 --- a/easytier-core/src/peers/mod.rs +++ b/easytier-core/src/peers/mod.rs @@ -1,5 +1,6 @@ pub(crate) mod acl; pub(crate) mod admission; +pub(crate) mod attached; pub(crate) mod conn; pub mod context; pub mod credential_manager; @@ -23,35 +24,134 @@ mod tests; use crate::packet::ZCPacket; use tokio::sync::mpsc::error::{SendError, TryRecvError, TrySendError}; -pub type PacketRecvChan = tokio::sync::mpsc::Sender; -pub type PacketRecvChanReceiver = tokio::sync::mpsc::Receiver; +use self::conn::peer_conn::PeerConnId; +use crate::config::PeerId; -pub fn create_packet_recv_chan() -> (PacketRecvChan, PacketRecvChanReceiver) { - tokio::sync::mpsc::channel(128) +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum PeerConnectionOrigin { + Network, + Attached, } -pub(crate) async fn send_packet_to_chan( +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum PeerPacketIngress { + Local, + Peer { + peer_id: PeerId, + conn_id: PeerConnId, + origin: PeerConnectionOrigin, + }, +} + +impl PeerPacketIngress { + pub(crate) fn is_attached(self) -> bool { + matches!( + self, + Self::Peer { + origin: PeerConnectionOrigin::Attached, + .. + } + ) + } + + pub(crate) fn peer_connection(self) -> Option<(PeerId, PeerConnId)> { + match self { + Self::Local => None, + Self::Peer { + peer_id, conn_id, .. + } => Some((peer_id, conn_id)), + } + } +} + +#[derive(Debug)] +pub(crate) struct PeerPacketEnvelope { + packet: ZCPacket, + ingress: PeerPacketIngress, +} + +impl PeerPacketEnvelope { + fn local(packet: ZCPacket) -> Self { + Self { + packet, + ingress: PeerPacketIngress::Local, + } + } + + fn from_peer(packet: ZCPacket, ingress: PeerPacketIngress) -> Self { + debug_assert!(matches!(ingress, PeerPacketIngress::Peer { .. })); + Self { packet, ingress } + } + + pub(crate) fn into_parts(self) -> (ZCPacket, PeerPacketIngress) { + (self.packet, self.ingress) + } +} + +#[derive(Clone)] +pub struct PacketRecvChan(tokio::sync::mpsc::Sender); + +impl PacketRecvChan { + /// Injects a packet produced by a local subsystem. Local injection never carries + /// the privileges of an attached peer connection. + pub async fn send(&self, packet: ZCPacket) -> Result<(), SendError> { + self.0 + .send(PeerPacketEnvelope::local(packet)) + .await + .map_err(|error| SendError(error.0.packet)) + } +} + +pub struct PacketRecvChanReceiver(tokio::sync::mpsc::Receiver); + +impl PacketRecvChanReceiver { + pub async fn recv(&mut self) -> Option { + self.0.recv().await.map(|envelope| envelope.packet) + } +} + +pub fn create_packet_recv_chan() -> (PacketRecvChan, PacketRecvChanReceiver) { + let (sender, receiver) = tokio::sync::mpsc::channel(128); + (PacketRecvChan(sender), PacketRecvChanReceiver(receiver)) +} + +pub(crate) async fn send_peer_packet_to_chan( sender: &PacketRecvChan, packet: ZCPacket, + ingress: PeerPacketIngress, ) -> Result<(), SendError> { - match sender.try_send(packet) { + let envelope = PeerPacketEnvelope::from_peer(packet, ingress); + match sender.0.try_send(envelope) { Ok(()) => Ok(()), - Err(TrySendError::Full(packet)) => sender.send(packet).await, - Err(TrySendError::Closed(packet)) => Err(SendError(packet)), + Err(TrySendError::Full(envelope)) => sender + .0 + .send(envelope) + .await + .map_err(|error| SendError(error.0.packet)), + Err(TrySendError::Closed(envelope)) => Err(SendError(envelope.packet)), + } +} + +pub(crate) async fn recv_packet_envelope_from_chan( + packet_recv_chan_receiver: &mut PacketRecvChanReceiver, +) -> Result { + match packet_recv_chan_receiver.0.try_recv() { + Ok(packet) => Ok(packet), + Err(TryRecvError::Empty) => packet_recv_chan_receiver + .0 + .recv() + .await + .ok_or(anyhow::anyhow!("recv_packet_from_chan failed")), + Err(TryRecvError::Disconnected) => Err(anyhow::anyhow!("recv_packet_from_chan failed")), } } pub async fn recv_packet_from_chan( packet_recv_chan_receiver: &mut PacketRecvChanReceiver, ) -> Result { - match packet_recv_chan_receiver.try_recv() { - Ok(packet) => Ok(packet), - Err(TryRecvError::Empty) => packet_recv_chan_receiver - .recv() - .await - .ok_or(anyhow::anyhow!("recv_packet_from_chan failed")), - Err(TryRecvError::Disconnected) => Err(anyhow::anyhow!("recv_packet_from_chan failed")), - } + recv_packet_envelope_from_chan(packet_recv_chan_receiver) + .await + .map(|envelope| envelope.packet) } #[async_trait::async_trait] diff --git a/easytier-core/src/peers/peer_manager.rs b/easytier-core/src/peers/peer_manager.rs index 061e9fb8..8006bcdf 100644 --- a/easytier-core/src/peers/peer_manager.rs +++ b/easytier-core/src/peers/peer_manager.rs @@ -21,9 +21,14 @@ use tokio::task::JoinSet; use url::Url; use crate::{ - config::peers::{HostRoutingPolicy, PeerRuntimeConfig, PeerRuntimeSnapshot}, - config::runtime::CoreRuntimeConfigStore, - config::{P2pPolicyFlags, PeerId, ProxyNetworkConfig}, + config::{ + P2pPolicyFlags, PeerId, ProxyNetworkConfig, + peers::{ + AclRuleConfig, HostRoutingPolicy, PeerGroupIdentity, PeerRuntimeConfig, + PeerRuntimeSnapshot, + }, + runtime::{CoreInstanceRuntimeConfig, CoreRuntimeConfigStore}, + }, events::CoreEventSink, foundation::task::ExternalTaskSignal, host::packet::{HostPacket, HostPacketSender}, @@ -43,7 +48,8 @@ use crate::{ }; use super::{ - BoxNicPacketFilter, BoxPeerPacketFilter, PacketRecvChanReceiver, PeerPacketFilter, + BoxNicPacketFilter, BoxPeerPacketFilter, PacketRecvChanReceiver, PeerConnectionOrigin, + PeerPacketFilter, PeerPacketIngress, acl::AclFilter, conn::{ peer_conn::{PeerConn, PeerConnId}, @@ -61,7 +67,7 @@ use super::{ peer_center::instance::PeerCenterPeerManagerTrait, peer_rpc::{PeerRpcManager, PeerRpcManagerTransport}, public_ipv6::{CorePublicIpv6Runtime, PublicIpv6Runtime}, - recv_packet_from_chan, + recv_packet_envelope_from_chan, relay_peer_map::RelayPeerMap, route::{ ArcRoute, DisabledRoute, ForeignNetworkRouteInfoMap, NextHopPolicy, Route, RouteInterface, @@ -238,6 +244,46 @@ impl PortablePeerManagerConfig { } } +fn matching_group_memberships( + declarations: &[PeerGroupIdentity], + configured_groups: &[String], +) -> Vec { + let configured = configured_groups.iter().collect::>(); + declarations + .iter() + .filter(|declaration| configured.contains(&declaration.group_name)) + .cloned() + .collect() +} + +fn first_missing_group<'a>( + memberships: &[PeerGroupIdentity], + configured_groups: &'a [String], +) -> Option<&'a str> { + let resolved = memberships + .iter() + .map(|membership| membership.group_name.as_str()) + .collect::>(); + configured_groups + .iter() + .map(String::as_str) + .find(|group| !resolved.contains(group)) +} + +fn retain_runtime_owned_peer_state( + current: &CoreInstanceRuntimeConfig, + next: &mut CoreInstanceRuntimeConfig, + peer_id: PeerId, +) { + let next_peer = Arc::make_mut(&mut next.peer); + next_peer.runtime.core.node.peer_id = Some(peer_id); + next_peer.runtime.core.node.instance_id = current.peer.runtime.core.node.instance_id; + next_peer.runtime.stun_info = current.peer.runtime.stun_info.clone(); + if current.services.dhcp_ipv4 && next.services.dhcp_ipv4 { + next_peer.runtime.core.routes.ipv4 = current.peer.runtime.core.routes.ipv4.clone(); + } +} + fn validate_portable_routes(routes: &crate::config::RouteConfig) -> anyhow::Result<()> { if !routes.advertised_routes.is_empty() { anyhow::bail!("portable peer manager does not support advertised routes yet"); @@ -724,6 +770,8 @@ pub struct PeerManagerCore { exit_nodes: Arc>>, acl_filter: Arc, context: Arc, + runtime_config: CoreRuntimeConfigStore, + runtime_config_update: Mutex<()>, is_secure_mode_enabled: bool, route: ArcRoute, traffic_metrics: Arc, @@ -771,6 +819,7 @@ impl PeerManagerCore { credential_storage: Option>, foreign_rpc_registrar: Arc, ) -> anyhow::Result { + let initial_acl = runtime_config.snapshot().services.acl.build()?; let runtime = &mut config.snapshot.runtime; let flags = &config.snapshot.flags; let network_name = runtime.network_identity.network_name.clone(); @@ -873,7 +922,7 @@ impl PeerManagerCore { credential_storage, }, )); - Ok(Self::assemble( + let peer_manager = Self::assemble( config.route_algo, my_peer_id, context, @@ -885,7 +934,9 @@ impl PeerManagerCore { config.exit_nodes, config.foreign_context_default_flags, foreign_rpc_registrar, - )) + ); + peer_manager.reload_acl(initial_acl.as_ref()); + Ok(peer_manager) } #[allow(clippy::too_many_arguments)] @@ -1086,6 +1137,8 @@ impl PeerManagerCore { data_compress_algo, exit_nodes, acl_filter, + runtime_config: core_context.runtime_config_store(), + runtime_config_update: Mutex::new(()), context: core_context, is_secure_mode_enabled, route, @@ -1103,6 +1156,64 @@ impl PeerManagerCore { self.context.credential_manager() } + pub(crate) fn register_ephemeral_credential( + &self, + public_key: [u8; 32], + groups: Vec, + allow_relay: bool, + allowed_proxy_cidrs: Vec, + reusable: bool, + ) -> anyhow::Result { + let credential_id = self + .credential_manager() + .register_ephemeral_credential( + public_key, + groups, + allow_relay, + allowed_proxy_cidrs, + reusable, + ) + .map_err(anyhow::Error::msg)?; + self.notify_credential_changed(); + Ok(credential_id) + } + + pub(crate) async fn update_ephemeral_credential_groups( + &self, + credential_id: uuid::Uuid, + groups: Vec, + ) -> Option { + let changed = self + .credential_manager() + .update_ephemeral_credential_groups(credential_id, groups)?; + if changed { + self.notify_credential_changed(); + self.route.refresh_acl_groups().await; + } + Some(changed) + } + + pub(crate) fn revoke_ephemeral_credential(&self, credential_id: uuid::Uuid) -> bool { + let revoked = self + .credential_manager() + .revoke_ephemeral_credential(credential_id); + if revoked { + self.notify_credential_changed(); + } + revoked + } + + pub(crate) async fn revoke_ephemeral_credential_and_refresh( + &self, + credential_id: uuid::Uuid, + ) -> bool { + let revoked = self.revoke_ephemeral_credential(credential_id); + if revoked { + self.route.refresh_acl_groups().await; + } + revoked + } + pub fn stats_manager(&self) -> Arc { self.context.stats_manager() } @@ -1219,6 +1330,122 @@ impl PeerManagerCore { self.acl_filter.clone() } + pub(crate) async fn update_runtime_config( + &self, + config: CoreInstanceRuntimeConfig, + ) -> anyhow::Result> { + let _update = self.runtime_config_update.lock().await; + let current = self.runtime_config.snapshot(); + let refresh_acl_groups = current.peer.peer_group_memberships + != config.peer.peer_group_memberships + || current.peer.acl_group_declarations != config.peer.acl_group_declarations; + let reload_acl = current.services.acl != config.services.acl; + let next_acl = reload_acl.then(|| config.services.acl.clone()); + let next_built_acl = next_acl.as_ref().map(AclRuleConfig::build).transpose()?; + + self.set_avoid_relay_data_preference(config.peer.avoid_relay_data_preference); + let published = self + .runtime_config + .replace_with_current(config, |current, next| { + retain_runtime_owned_peer_state(current, next, self.my_peer_id); + }); + if let Some(acl) = next_built_acl.as_ref() { + self.reload_acl(acl.as_ref()); + } + if refresh_acl_groups { + self.route.refresh_acl_groups().await; + } + Ok(published) + } + + /// Keeps this manager's ACL rules and group assignments aligned with a + /// network policy source. The subscription is owned by this manager and is + /// stopped with its other runtime tasks. + pub(crate) async fn follow_network_policy( + self: &Arc, + source: CoreRuntimeConfigStore, + configured_groups: Vec, + ) -> anyhow::Result<()> { + let mut peer_changes = source.subscribe_peer_runtime_changes(); + let mut service_changes = source.subscribe_service_runtime_changes(); + let configured_groups: Arc<[String]> = configured_groups.into(); + self.apply_network_policy(&source, &configured_groups) + .await?; + + let peer_manager = Arc::downgrade(self); + self.tasks.lock().await.spawn(async move { + loop { + let changed = tokio::select! { + changed = peer_changes.changed() => changed, + changed = service_changes.changed() => changed, + }; + if changed.is_err() { + return; + } + let _ = peer_changes.borrow_and_update(); + let _ = service_changes.borrow_and_update(); + let Some(peer_manager) = peer_manager.upgrade() else { + return; + }; + if let Err(error) = peer_manager + .apply_network_policy(&source, &configured_groups) + .await + { + tracing::warn!( + ?error, + peer_id = peer_manager.my_peer_id, + "failed to apply peer manager network policy" + ); + } + } + }); + Ok(()) + } + + async fn apply_network_policy( + &self, + source: &CoreRuntimeConfigStore, + configured_groups: &[String], + ) -> anyhow::Result<()> { + let source = source.snapshot(); + let credential_peer = self.context.feature_flags().is_credential_peer; + let acl = if credential_peer { + source.services.acl.for_credential_peer() + } else { + source.services.acl.clone() + }; + let (declarations, memberships, missing_group) = if credential_peer { + (Vec::new(), Vec::new(), None) + } else { + let declarations = source.peer.acl_group_declarations.clone(); + let memberships = matching_group_memberships(&declarations, configured_groups); + let missing_group = first_missing_group(&memberships, configured_groups); + (declarations, memberships, missing_group) + }; + let current = self.runtime_config.snapshot(); + if current.services.acl == acl + && current.peer.acl_group_declarations == declarations + && current.peer.peer_group_memberships == memberships + { + return Ok(()); + } + + let mut next = current.as_ref().clone(); + next.services.acl = acl; + let peer = Arc::make_mut(&mut next.peer); + peer.acl_group_declarations = declarations; + peer.peer_group_memberships = memberships; + self.update_runtime_config(next).await?; + if let Some(group) = missing_group { + tracing::warn!( + peer_id = self.my_peer_id, + group, + "configured peer ACL group is no longer declared" + ); + } + Ok(()) + } + pub fn network_name(&self) -> &str { &self.network_name } @@ -1264,7 +1491,6 @@ impl PeerManagerCore { pub fn get_peer_session_store(&self) -> Arc { self.peer_session_store.clone() } - pub(crate) fn get_nic_channel(&self) -> HostPacketSender { self.nic_channel.clone() } @@ -1353,6 +1579,15 @@ impl PeerManagerCore { .await } + pub(crate) async fn add_attached_ring_client_tunnel( + &self, + tunnel: Box, + ) -> Result<(PeerId, PeerConnId), Error> { + self.peer_connection_admission + .add_client_tunnel_with_origin(tunnel, true, None, PeerConnectionOrigin::Attached) + .await + } + pub async fn add_tunnel_as_server( &self, tunnel: Box, @@ -1363,6 +1598,15 @@ impl PeerManagerCore { .await } + pub(crate) async fn add_attached_ring_tunnel_as_server( + &self, + tunnel: Box, + ) -> Result<(PeerId, PeerConnId), Error> { + self.peer_connection_admission + .add_tunnel_as_server_with_origin(tunnel, true, PeerConnectionOrigin::Attached) + .await + } + pub async fn add_packet_process_pipeline(&self, pipeline: BoxPeerPacketFilter) { // newest pipeline will be executed first append_peer_pipeline( @@ -1501,11 +1745,7 @@ impl PeerManagerCore { *self.exit_nodes.write().await = exit_nodes; } - pub(crate) fn reload_acl(&self, acl: Option<&crate::proto::acl::Acl>) { - // ACL rule effects are staged separately from configuration publication. - // Keep the submitted group snapshot unchanged so CoreInstance can detect - // the group change and refresh route trust state when the complete - // runtime configuration is published. + fn reload_acl(&self, acl: Option<&crate::proto::acl::Acl>) { self.acl_filter.reload_rules(acl); } @@ -1532,6 +1772,12 @@ impl PeerManagerCore { pub(crate) async fn clear_resources(&self) { self.stop().await; + self.foreign_network_client.stop(); + self.peers.clear_resources().await; + self.foreign_network_client + .get_peer_map() + .clear_resources() + .await; self.peer_packet_process_pipeline .store(Arc::new(Vec::new())); self.nic_packet_process_pipeline.store(Arc::new(Vec::new())); @@ -1723,19 +1969,43 @@ impl PeerConnectionAdmission { is_directly_connected: bool, peer_id_hint: Option, ) -> Result<(PeerId, PeerConnId), Error> { - let mut peer = PeerConn::new_with_peer_id_hint( + self.add_client_tunnel_with_origin( + tunnel, + is_directly_connected, + peer_id_hint, + PeerConnectionOrigin::Network, + ) + .await + } + + async fn add_client_tunnel_with_origin( + &self, + tunnel: Box, + is_directly_connected: bool, + peer_id_hint: Option, + origin: PeerConnectionOrigin, + ) -> Result<(PeerId, PeerConnId), Error> { + let mut peer = PeerConn::new_with_peer_id_hint_and_origin( self.my_peer_id, self.context.clone(), tunnel, peer_id_hint, self.peer_session_store.clone(), + origin, ); peer.set_is_hole_punched(!is_directly_connected); peer.do_handshake_as_client().await?; let conn_id = peer.get_conn_id(); let peer_id = peer.get_peer_id(); let local_identity = self.context.network_identity(); - if peer.get_network_identity().network_name == local_identity.network_name { + let is_local_network = + peer.get_network_identity().network_name == local_identity.network_name; + if origin == PeerConnectionOrigin::Attached && !is_local_network { + return Err(Error::SecretKeyError( + "attached ring peer must belong to the local network".to_string(), + )); + } + if is_local_network { let local_secure_mode = self .context .secure_mode() @@ -1767,15 +2037,32 @@ impl PeerConnectionAdmission { tunnel: Box, is_directly_connected: bool, ) -> Result<(), Error> { + self.add_tunnel_as_server_with_origin( + tunnel, + is_directly_connected, + PeerConnectionOrigin::Network, + ) + .await + .map(|_| ()) + } + + async fn add_tunnel_as_server_with_origin( + &self, + tunnel: Box, + is_directly_connected: bool, + origin: PeerConnectionOrigin, + ) -> Result<(PeerId, PeerConnId), Error> { tracing::info!("add tunnel as server start"); let resolved_remote_addr = tunnel.info().and_then(|info| info.resolved_remote_addr); check_resolved_remote_addr_not_from_virtual_network(&self.context, resolved_remote_addr)?; - let mut conn = PeerConn::new( + let mut conn = PeerConn::new_with_peer_id_hint_and_origin( self.my_peer_id, self.context.clone(), tunnel, + None, self.peer_session_store.clone(), + origin, ); let mut reserved_peer_id_network_name = None; let handshake_ret = conn @@ -1821,6 +2108,14 @@ impl PeerConnectionAdmission { let peer_network_name = peer_identity.network_name.clone(); let local_identity = self.context.network_identity(); let is_local_network = peer_network_name == local_identity.network_name; + if origin == PeerConnectionOrigin::Attached && !is_local_network { + self.release_reserved_peer_id(&peer_network_name); + return Err(Error::SecretKeyError( + "attached ring peer must belong to the local network".to_string(), + )); + } + let peer_id = conn.get_peer_id(); + let conn_id = conn.get_conn_id(); let trusted_foreign_credential = matches!(conn.get_peer_identity_type(), PeerIdentityType::Credential) && self @@ -1874,7 +2169,7 @@ impl PeerConnectionAdmission { self.release_reserved_peer_id(&peer_network_name); tracing::info!("add tunnel as server done"); - Ok(()) + Ok((peer_id, conn_id)) } } @@ -2721,29 +3016,51 @@ impl PeerPacketRouter { pub async fn run(mut self) { tracing::trace!("start_peer_recv"); - while let Ok(ret) = recv_packet_from_chan(&mut self.packet_recv).await { + while let Ok(envelope) = recv_packet_envelope_from_chan(&mut self.packet_recv).await { + let (ret, ingress) = envelope.into_parts(); let disable_relay_data = self.context.disable_relay_data(); + let destination_is_attached = ret + .peer_manager_header() + .is_some_and(|header| self.peers.has_direct_attached_peer(header.to_peer_id.get())); + let drop_foreign_relay_data = disable_relay_data + && is_relay_data_zc_packet(&ret) + && !ingress.is_attached() + && !destination_is_attached; let Err(ret) = try_handle_foreign_network_packet( ret, self.my_peer_id, &self.peers, &self.foreign_network_manager, self.stats_mgr.as_ref(), - disable_relay_data, + drop_foreign_relay_data, ) .await else { continue; }; - self.handle_packet(ret, disable_relay_data).await; + self.handle_packet(ret, disable_relay_data, ingress).await; } panic!("done_peer_recv"); } - async fn handle_packet(&self, mut ret: ZCPacket, disable_relay_data: bool) { + async fn handle_packet( + &self, + mut ret: ZCPacket, + disable_relay_data: bool, + ingress: PeerPacketIngress, + ) { let buf_len = ret.buf_len(); - let is_relay_data_packet = is_relay_data_zc_packet(&ret); + let destination_is_attached = ret + .peer_manager_header() + .is_some_and(|header| self.peers.has_direct_attached_peer(header.to_peer_id.get())); + let drop_relay_data = should_drop_relay_data( + disable_relay_data, + &ret, + self.my_peer_id, + ingress, + destination_is_attached, + ); let Some(hdr) = ret.mut_peer_manager_header() else { tracing::warn!(?ret, "invalid packet, skip"); return; @@ -2755,11 +3072,13 @@ impl PeerPacketRouter { let packet_type = hdr.packet_type; let is_encrypted = hdr.is_encrypted(); if to_peer_id != self.my_peer_id { - if disable_relay_data && is_relay_data_packet { + if drop_relay_data { + let ingress_connection = ingress.peer_connection(); tracing::debug!( ?from_peer_id, ?to_peer_id, packet_type, + ?ingress_connection, "drop forwarded relay data while relay data is disabled" ); return; @@ -2938,13 +3257,30 @@ pub(crate) fn is_relay_data_zc_packet(packet: &ZCPacket) -> bool { is_relay_data_packet(hdr.packet_type) } +fn should_drop_relay_data( + disable_relay_data: bool, + packet: &ZCPacket, + my_peer_id: PeerId, + ingress: PeerPacketIngress, + destination_is_attached: bool, +) -> bool { + if !disable_relay_data || !is_relay_data_zc_packet(packet) { + return false; + } + + let is_forwarded = packet + .peer_manager_header() + .is_some_and(|header| header.to_peer_id.get() != my_peer_id); + is_forwarded && !ingress.is_attached() && !destination_is_attached +} + pub(crate) async fn try_handle_foreign_network_packet( mut packet: ZCPacket, my_peer_id: PeerId, peer_map: &PeerMap, foreign_network_manager: &ForeignNetworkManager, stats_manager: &StatsManager, - disable_relay_data: bool, + drop_relay_data: bool, ) -> Result<(), ZCPacket> { let pm_header = packet.peer_manager_header().unwrap(); if pm_header.packet_type != PacketType::ForeignNetworkPacket as u8 { @@ -2954,7 +3290,7 @@ pub(crate) async fn try_handle_foreign_network_packet( let from_peer_id = pm_header.from_peer_id.get(); let to_peer_id = pm_header.to_peer_id.get(); - if disable_relay_data && is_relay_data_zc_packet(&packet) { + if drop_relay_data { tracing::debug!( ?from_peer_id, ?to_peer_id, @@ -3413,6 +3749,35 @@ mod tests { .await ); } + #[tokio::test] + async fn runtime_updates_retain_manager_owned_peer_identity() { + let core = build_portable_for_test(portable_runtime_config("portable-net")).unwrap(); + let current = core.runtime_config.snapshot(); + let expected_peer_id = core.my_peer_id(); + let expected_instance_id = current.peer.runtime.core.node.instance_id; + let mut next = current.as_ref().clone(); + let next_peer = Arc::make_mut(&mut next.peer); + next_peer.runtime.core.node.peer_id = Some(expected_peer_id.wrapping_add(1)); + next_peer.runtime.core.node.instance_id = Some([1; 16]); + let submitted = next.peer.clone(); + + let published = core.update_runtime_config(next).await.unwrap(); + + assert_eq!( + published.peer.runtime.core.node.peer_id, + Some(expected_peer_id) + ); + assert_eq!( + published.peer.runtime.core.node.instance_id, + expected_instance_id + ); + assert_eq!( + submitted.runtime.core.node.peer_id, + Some(expected_peer_id.wrapping_add(1)) + ); + assert_eq!(submitted.runtime.core.node.instance_id, Some([1; 16])); + core.clear_resources().await; + } #[cfg(not(feature = "zstd"))] #[tokio::test] @@ -3557,6 +3922,83 @@ mod tests { core.clear_resources().await; } + #[tokio::test] + async fn credential_peer_policy_sync_excludes_group_material() { + let mut runtime = portable_runtime_config("portable-net"); + runtime.network_identity.network_secret = None; + runtime.network_identity.network_secret_digest = None; + runtime.secure_mode = Some(credential_secure_mode()); + let core = Arc::new(build_portable_for_test(runtime).unwrap()); + + let acl = crate::proto::acl::Acl { + acl_v1: Some(crate::proto::acl::AclV1 { + chains: Vec::new(), + group: Some(crate::proto::acl::GroupInfo { + declares: vec![crate::proto::acl::GroupIdentity { + group_name: "ops".to_owned(), + group_secret: "ops-secret".to_owned(), + }], + members: vec!["ops".to_owned()], + }), + }), + }; + let mut services = CoreRuntimeConfig::default(); + services.acl.acl = Some(acl.clone()); + let mut source_peer = core.runtime_config.snapshot().peer.as_ref().clone(); + source_peer.set_acl_groups(Some(&acl)); + let source = CoreRuntimeConfigStore::new(services, Arc::new(source_peer)); + + core.follow_network_policy(source.clone(), vec!["ops".to_owned()]) + .await + .unwrap(); + + let applied = core.runtime_config.snapshot(); + assert!(applied.peer.acl_group_declarations.is_empty()); + assert!(applied.peer.peer_group_memberships.is_empty()); + assert!( + applied + .services + .acl + .acl + .as_ref() + .unwrap() + .acl_v1 + .as_ref() + .unwrap() + .group + .is_none() + ); + + source.update_services(|services| { + services.acl.tcp_whitelist = vec!["22".to_owned()]; + }); + tokio::time::timeout(Duration::from_secs(10), async { + loop { + let applied = core.runtime_config.snapshot(); + if applied.services.acl.tcp_whitelist == ["22"] { + assert!( + applied + .services + .acl + .acl + .as_ref() + .unwrap() + .acl_v1 + .as_ref() + .unwrap() + .group + .is_none() + ); + return; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + core.clear_resources().await; + } + #[tokio::test] async fn node_snapshot_exposes_normalized_runtime_state() { let instance_id = uuid::Uuid::from_u128(0x00112233445566778899aabbccddeeff); @@ -3966,6 +4408,60 @@ mod tests { assert!(is_relay_data_zc_packet(&foreign_data_packet)); } + fn data_packet(from_peer_id: PeerId, to_peer_id: PeerId) -> ZCPacket { + let mut packet = ZCPacket::new_with_payload(b"data"); + packet.fill_peer_manager_hdr(from_peer_id, to_peer_id, PacketType::Data as u8); + packet + } + + #[test] + fn forged_attached_source_header_does_not_bypass_relay_disable() { + let packet = data_packet(77, 3); + let network_ingress = PeerPacketIngress::Peer { + peer_id: 2, + conn_id: PeerConnId::new_v4(), + origin: PeerConnectionOrigin::Network, + }; + + assert!(should_drop_relay_data( + true, + &packet, + 1, + network_ingress, + false, + )); + } + + #[test] + fn attached_ingress_and_destination_bypass_relay_disable() { + let packet = data_packet(2, 3); + let attached_ingress = PeerPacketIngress::Peer { + peer_id: 2, + conn_id: PeerConnId::new_v4(), + origin: PeerConnectionOrigin::Attached, + }; + let network_ingress = PeerPacketIngress::Peer { + peer_id: 2, + conn_id: PeerConnId::new_v4(), + origin: PeerConnectionOrigin::Network, + }; + + assert!(!should_drop_relay_data( + true, + &packet, + 1, + attached_ingress, + false, + )); + assert!(!should_drop_relay_data( + true, + &packet, + 1, + network_ingress, + true, + )); + } + fn route_with_ipv4( peer_id: u32, ipv4_addr: Option, diff --git a/easytier-core/src/peers/route/peer_ospf_route.rs b/easytier-core/src/peers/route/peer_ospf_route.rs index 7029c4f4..253394d0 100644 --- a/easytier-core/src/peers/route/peer_ospf_route.rs +++ b/easytier-core/src/peers/route/peer_ospf_route.rs @@ -793,7 +793,6 @@ pub fn new_updated_self_route_peer_info( proxy_cidrs: context .proxy_cidrs() .into_iter() - .chain(context.vpn_portal_cidr()) .map(|x| x.to_string()) .collect(), hostname: Some(context.hostname()), diff --git a/easytier-core/src/peers/tests.rs b/easytier-core/src/peers/tests.rs index d2d6256d..4da0107d 100644 --- a/easytier-core/src/peers/tests.rs +++ b/easytier-core/src/peers/tests.rs @@ -8,6 +8,7 @@ use crate::foundation::time::{Duration, timeout}; use crate::{ packet::{PacketType, ZCPacket}, peers::{ + PeerConnectionOrigin, PeerPacketIngress, conn::{ peer_conn::{PeerConn, PeerConnId}, peer_map::PeerMap, @@ -16,6 +17,7 @@ use crate::{ context::NetworkIdentity, create_packet_recv_chan, error::Error, + recv_packet_envelope_from_chan, test_support::NoopPeerContext, }, tunnel::ring::create_ring_tunnel_pair, @@ -124,6 +126,20 @@ async fn peer_conn_handshake_matches_plaintext_secret_identity() { assert!(server.matches_local_network_secret()); } +#[tokio::test] +async fn local_packet_channel_injection_has_no_attached_privilege() { + let (sender, mut receiver) = create_packet_recv_chan(); + let mut packet = ZCPacket::new_with_payload(b"local"); + packet.fill_peer_manager_hdr(77, 3, PacketType::Data as u8); + + sender.send(packet).await.unwrap(); + + let envelope = recv_packet_envelope_from_chan(&mut receiver).await.unwrap(); + let (_, ingress) = envelope.into_parts(); + assert_eq!(ingress, PeerPacketIngress::Local); + assert!(!ingress.is_attached()); +} + #[tokio::test] async fn peer_map_forwards_packet_over_memory_tunnel() { let peer_session_store = Arc::new(PeerSessionStore::new()); @@ -165,6 +181,70 @@ async fn peer_map_forwards_packet_over_memory_tunnel() { assert_eq!(received.payload(), b"hello"); } +#[tokio::test] +async fn peer_channel_uses_admission_origin_instead_of_packet_header() { + let peer_session_store = Arc::new(PeerSessionStore::new()); + let (client_tunnel, server_tunnel) = create_ring_tunnel_pair(); + let client_ctx = Arc::new(NoopPeerContext::default()); + let server_ctx = Arc::new(NoopPeerContext::default()); + + let mut client_conn = PeerConn::new( + 1, + client_ctx.clone(), + client_tunnel, + peer_session_store.clone(), + ); + let mut server_conn = PeerConn::new_with_peer_id_hint_and_origin( + 2, + server_ctx.clone(), + server_tunnel, + None, + peer_session_store, + PeerConnectionOrigin::Attached, + ); + let (client_ret, server_ret) = tokio::join!( + client_conn.do_handshake_as_client(), + server_conn.do_handshake_as_server() + ); + client_ret.unwrap(); + server_ret.unwrap(); + server_conn.set_is_hole_punched(false); + let server_conn_id = server_conn.get_conn_id(); + + let (client_tx, _client_rx) = create_packet_recv_chan(); + let (server_tx, mut server_rx) = create_packet_recv_chan(); + let client_map = PeerMap::new(client_tx, client_ctx, 1); + let server_map = PeerMap::new(server_tx, server_ctx, 2); + client_map.add_new_peer_conn(client_conn).await.unwrap(); + server_map.add_new_peer_conn(server_conn).await.unwrap(); + assert!(server_map.has_direct_attached_peer(1)); + + let mut packet = ZCPacket::new_with_payload(b"forged source"); + packet.fill_peer_manager_hdr(77, 3, PacketType::Data as u8); + client_map.send_msg_directly(packet, 2).await.unwrap(); + + let envelope = timeout( + Duration::from_secs(1), + recv_packet_envelope_from_chan(&mut server_rx), + ) + .await + .unwrap() + .unwrap(); + let (received, ingress) = envelope.into_parts(); + assert_eq!( + received.peer_manager_header().unwrap().from_peer_id.get(), + 77 + ); + assert_eq!( + ingress, + PeerPacketIngress::Peer { + peer_id: 1, + conn_id: server_conn_id, + origin: PeerConnectionOrigin::Attached, + } + ); +} + #[tokio::test] async fn peer_map_reselects_cached_connection_after_close() { let peer_session_store = Arc::new(PeerSessionStore::new()); diff --git a/easytier-gui/src-tauri/src/lib.rs b/easytier-gui/src-tauri/src/lib.rs index ce4bd86b..c6c169e0 100644 --- a/easytier-gui/src-tauri/src/lib.rs +++ b/easytier-gui/src-tauri/src/lib.rs @@ -6,10 +6,15 @@ mod elevate; use anyhow::Context; #[cfg(target_os = "android")] use easytier::instance::factory::subscribe_native_instance_event; +use easytier::proto::api::instance::{ + GetVpnPortalInfoRequest, InstanceIdentifier, VpnPortalInfo, VpnPortalRpc, + VpnPortalRpcClientFactory, instance_identifier, +}; use easytier::proto::api::manage::{ CollectNetworkInfoResponse, ValidateConfigResponse, WebClientService, WebClientServiceClientFactory, }; +use easytier::proto::rpc_types::controller::BaseController; use easytier::web_client::{self, WebClient}; use easytier::{ common::config::{NetworkConfig, NetworkConfigExt}, @@ -151,6 +156,30 @@ async fn collect_network_info( .map_err(|e| e.to_string()) } +#[tauri::command] +async fn get_vpn_portal_info(instance_id: String) -> Result, String> { + let instance_id = instance_id + .parse::() + .map_err(|e| e.to_string())?; + let client_manager = get_client_manager!()?; + let client = client_manager + .rpc_manager + .rpc_client() + .scoped_client::>(1, 1, "".to_string()); + let response = client + .get_vpn_portal_info( + BaseController::default(), + GetVpnPortalInfoRequest { + instance: Some(InstanceIdentifier { + selector: Some(instance_identifier::Selector::Id(instance_id.into())), + }), + }, + ) + .await + .map_err(|e| e.to_string())?; + Ok(response.vpn_portal_info) +} + #[tauri::command] async fn set_logging_level(level: String) -> Result<(), String> { get_client_manager!()? @@ -1393,6 +1422,7 @@ pub fn run_gui() -> std::process::ExitCode { generate_network_config, run_network_instance, collect_network_info, + get_vpn_portal_info, set_logging_level, set_tun_fd, easytier_version, diff --git a/easytier-gui/src/auto-imports.d.ts b/easytier-gui/src/auto-imports.d.ts index f1450cb5..571aff10 100644 --- a/easytier-gui/src/auto-imports.d.ts +++ b/easytier-gui/src/auto-imports.d.ts @@ -29,6 +29,7 @@ declare global { const getEasytierVersion: typeof import('./composables/backend')['getEasytierVersion'] const getNetworkMetas: typeof import('./composables/backend')['getNetworkMetas'] const getServiceStatus: typeof import('./composables/backend')['getServiceStatus'] + const getVpnPortalInfo: typeof import('./composables/backend')['getVpnPortalInfo'] const h: typeof import('vue')['h'] const initMobileVpnService: typeof import('./composables/mobile_vpn')['initMobileVpnService'] const initRpcConnection: typeof import('./composables/backend')['initRpcConnection'] @@ -156,6 +157,7 @@ declare module 'vue' { readonly getEasytierVersion: UnwrapRef readonly getNetworkMetas: UnwrapRef readonly getServiceStatus: UnwrapRef + readonly getVpnPortalInfo: UnwrapRef readonly h: UnwrapRef readonly initMobileVpnService: UnwrapRef readonly initRpcConnection: UnwrapRef diff --git a/easytier-gui/src/composables/backend.ts b/easytier-gui/src/composables/backend.ts index a16835f5..1a779aa7 100644 --- a/easytier-gui/src/composables/backend.ts +++ b/easytier-gui/src/composables/backend.ts @@ -66,6 +66,11 @@ export async function collectNetworkInfo(instanceId: string) { return await invoke('collect_network_info', { instanceId }) } +export async function getVpnPortalInfo(instanceId: string) { + const info = await invoke('get_vpn_portal_info', { instanceId }) + return info ? NetworkTypes.normalizeVpnPortalInfo(info) : undefined +} + export async function setLoggingLevel(level: string) { return await invoke('set_logging_level', { level }) } diff --git a/easytier-gui/src/modules/api.ts b/easytier-gui/src/modules/api.ts index d24c27e6..ca4a97ae 100644 --- a/easytier-gui/src/modules/api.ts +++ b/easytier-gui/src/modules/api.ts @@ -11,6 +11,9 @@ export class GUIRemoteClient implements Api.RemoteClient { async get_network_info(inst_id: string): Promise { return backend.collectNetworkInfo(inst_id).then(infos => infos.info?.map?.[inst_id]); } + async get_vpn_portal_info(inst_id: string): Promise { + return backend.getVpnPortalInfo(inst_id); + } async list_network_instance_ids(): Promise { return backend.listNetworkInstanceIds(); } diff --git a/easytier-proto/build/main.rs b/easytier-proto/build/main.rs index b8b75c27..8c9752c4 100644 --- a/easytier-proto/build/main.rs +++ b/easytier-proto/build/main.rs @@ -124,7 +124,12 @@ fn main() -> Result<(), Box> { .file_descriptor_set_path(&descriptor) .service_generator(Box::new(ServiceGenerator::default())) .btree_map(["."]) - .skip_debug([".common.Ipv4Addr", ".common.Ipv6Addr", ".common.UUID"]); + .skip_debug([ + ".common.Ipv4Addr", + ".common.Ipv6Addr", + ".common.UUID", + ".api.manage.VpnPortalConfig", + ]); config.compile_protos(&proto_files, &["proto/"])?; diff --git a/easytier-proto/proto/api_instance.proto b/easytier-proto/proto/api_instance.proto index c4512d32..823e632a 100644 --- a/easytier-proto/proto/api_instance.proto +++ b/easytier-proto/proto/api_instance.proto @@ -231,10 +231,32 @@ service MappedListenerManageRpc { returns (ListMappedListenerResponse); } +enum VpnPortalClientState { + VPN_PORTAL_CLIENT_STATE_UNSPECIFIED = 0; + VPN_PORTAL_CLIENT_STATE_OFFLINE = 1; + VPN_PORTAL_CLIENT_STATE_CONNECTING = 2; + VPN_PORTAL_CLIENT_STATE_ONLINE = 3; + VPN_PORTAL_CLIENT_STATE_ERROR = 4; +} + +message VpnPortalClientInfo { + string name = 1; + string virtual_ip = 2; + repeated string groups = 3; + VpnPortalClientState state = 4; + optional uint32 peer_id = 5; + optional string endpoint = 6; + optional string tunnel_ip = 7; + string client_config = 8; + optional string error = 9; +} + message VpnPortalInfo { string vpn_type = 1; - string client_config = 2; - repeated string connected_clients = 3; + string client_config = 2 [deprecated = true]; + repeated string connected_clients = 3 [deprecated = true]; + repeated VpnPortalClientInfo clients = 4; + optional string listener = 5; } message GetVpnPortalInfoRequest { InstanceIdentifier instance = 1; } diff --git a/easytier-proto/proto/api_manage.proto b/easytier-proto/proto/api_manage.proto index 84e79a0f..6f67a901 100644 --- a/easytier-proto/proto/api_manage.proto +++ b/easytier-proto/proto/api_manage.proto @@ -35,10 +35,10 @@ message NetworkConfig { repeated string proxy_cidrs = 11; - optional bool enable_vpn_portal = 12; - optional int32 vpn_portal_listen_port = 13; - optional string vpn_portal_client_network_addr = 14; - optional int32 vpn_portal_client_network_len = 15; + optional bool enable_vpn_portal = 12 [deprecated = true]; + optional int32 vpn_portal_listen_port = 13 [deprecated = true]; + optional string vpn_portal_client_network_addr = 14 [deprecated = true]; + optional int32 vpn_portal_client_network_len = 15 [deprecated = true]; optional bool advanced_settings = 16; @@ -103,6 +103,19 @@ message NetworkConfig { optional bool enable_udp_broadcast_relay = 66; optional uint32 socket_mark = 67; repeated NetworkPeerConfig peers = 68; + optional VpnPortalConfig vpn_portal_config = 69; +} + +message VpnPortalClientConfig { + string name = 1; + string virtual_ip = 2; + repeated string groups = 3; +} + +message VpnPortalConfig { + string wireguard_listen = 1; + optional string wireguard_private_key = 2; + repeated VpnPortalClientConfig clients = 3; } message NetworkPeerConfig { @@ -125,7 +138,7 @@ message MyNodeInfo { peer_rpc.GetIpListResponse ips = 4; common.StunInfo stun_info = 5; repeated common.Url listeners = 6; - optional string vpn_portal_cfg = 7; + optional string vpn_portal_cfg = 7 [deprecated = true]; uint32 peer_id = 8; } diff --git a/easytier-proto/src/api.rs b/easytier-proto/src/api.rs index fcb6d97b..efc552c7 100644 --- a/easytier-proto/src/api.rs +++ b/easytier-proto/src/api.rs @@ -334,6 +334,20 @@ pub mod manage { include!(concat!(env!("OUT_DIR"), "/api.manage.rs")); #[cfg(feature = "json-rpc")] include!(concat!(env!("OUT_DIR"), "/api.manage.serde.rs")); + + impl std::fmt::Debug for VpnPortalConfig { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("VpnPortalConfig") + .field("wireguard_listen", &self.wireguard_listen) + .field( + "wireguard_private_key", + &self.wireguard_private_key.as_ref().map(|_| ""), + ) + .field("clients", &self.clients) + .finish() + } + } } #[cfg(test)] @@ -352,6 +366,19 @@ mod tests { use crate::proto::rpc_types::error::Error; use crate::proto::rpc_types::handler::Handler; + #[test] + fn vpn_portal_debug_redacts_private_key() { + let config = super::manage::VpnPortalConfig { + wireguard_listen: "0.0.0.0:51820".to_owned(), + wireguard_private_key: Some("private-key-material".to_owned()), + clients: Vec::new(), + }; + + let debug = format!("{config:?}"); + assert!(debug.contains("")); + assert!(!debug.contains("private-key-material")); + } + #[derive(Clone, Default)] struct WebClientServiceJsonCallHandler; diff --git a/easytier-web/frontend-lib/scripts/test-network-config.mjs b/easytier-web/frontend-lib/scripts/test-network-config.mjs index a996427d..3e08faf3 100644 --- a/easytier-web/frontend-lib/scripts/test-network-config.mjs +++ b/easytier-web/frontend-lib/scripts/test-network-config.mjs @@ -20,12 +20,12 @@ const { DEFAULT_NETWORK_CONFIG, NetworkingMethod, normalizeNetworkConfig, + normalizeVpnPortalInfo, toBackendNetworkConfig, } = NetworkTypes const BOOLEAN_CONFIG_FIELDS = [ 'dhcp', - 'enable_vpn_portal', 'advanced_settings', 'latency_first', 'use_smoltcp', @@ -60,6 +60,13 @@ const BOOLEAN_CONFIG_FIELDS = [ 'disable_tcp_hole_punching', ] +const LEGACY_VPN_PORTAL_FIELDS = [ + 'enable_vpn_portal', + 'vpn_portal_listen_port', + 'vpn_portal_client_network_addr', + 'vpn_portal_client_network_len', +] + function readGeneratedNetworkConfigFields() { const source = ts.createSourceFile( generatedApiManagePath, @@ -121,10 +128,15 @@ function allFieldFixture() { }, ], proxy_cidrs: ['10.10.0.0/16', '192.168.2.0/24->10.99.0.0/24'], - enable_vpn_portal: true, - vpn_portal_listen_port: 23000, - vpn_portal_client_network_addr: '10.88.0.0', - vpn_portal_client_network_len: 24, + vpn_portal_config: { + wireguard_listen: '0.0.0.0:23000', + wireguard_private_key: 'portal-private-key', + clients: [{ + name: 'phone-a', + virtual_ip: '10.9.8.10', + groups: ['ops'], + }], + }, advanced_settings: true, listener_urls: ['tcp://0.0.0.0:12010', 'udp://0.0.0.0:12010'], latency_first: true, @@ -239,6 +251,7 @@ function allFieldFixture() { function assertFixtureCoversGeneratedFields() { const generatedFields = readGeneratedNetworkConfigFields() + .filter((field) => !LEGACY_VPN_PORTAL_FIELDS.includes(field)) const fixtureFields = new Set(Object.keys(allFieldFixture())) const missing = generatedFields.filter((field) => !fixtureFields.has(field)) @@ -258,7 +271,8 @@ function assertFullFieldRoundTrip() { const backend = toBackendNetworkConfig(normalized) expectNoCamelCaseKeys(backend) - for (const field of readGeneratedNetworkConfigFields()) { + for (const field of readGeneratedNetworkConfigFields() + .filter((field) => !LEGACY_VPN_PORTAL_FIELDS.includes(field))) { assert.ok(field in backend, `backend JSON should include fixture field ${field}`) } @@ -269,6 +283,9 @@ function assertFullFieldRoundTrip() { assert.deepEqual(backend.peers[1], { uri: 'udp://peer-b:11010' }) assert.equal(backend.data_compress_algo, 'Zstd') assert.equal(backend.instance_recv_bps_limit, '9007199254740993') + assert.equal(backend.vpn_portal_config.wireguard_listen, '0.0.0.0:23000') + assert.equal(backend.vpn_portal_config.clients[0].name, 'phone-a') + assert.deepEqual(backend.vpn_portal_config.clients[0].groups, ['ops']) assert.equal(backend.secure_mode.enabled, true) assert.equal(backend.secure_mode.local_private_key, 'private-key') assert.equal(backend.acl.acl_v1.chains[0].chain_type, 'Forward') @@ -279,6 +296,38 @@ function assertFullFieldRoundTrip() { assert.equal(backend.socket_mark, 1234) } +function assertLegacyVpnPortalFieldsReachBackendValidation() { + const backend = toBackendNetworkConfig({ + ...DEFAULT_NETWORK_CONFIG(), + enable_vpn_portal: true, + vpn_portal_listen_port: 22022, + vpn_portal_client_network_addr: '10.88.0.0', + vpn_portal_client_network_len: 24, + }) + + assert.equal(backend.enable_vpn_portal, true) + assert.equal(backend.vpn_portal_listen_port, 22022) + assert.equal(backend.vpn_portal_client_network_addr, '10.88.0.0') + assert.equal(backend.vpn_portal_client_network_len, 24) +} + +function assertVpnPortalRpcJsonNormalization() { + const info = normalizeVpnPortalInfo({ + vpn_type: 'wireguard', + listener: '0.0.0.0:22022', + clients: [{ + name: 'phone-a', + virtual_ip: '10.9.8.10', + groups: ['ops'], + state: 'VPN_PORTAL_CLIENT_STATE_ONLINE', + client_config: '[Interface]', + }], + }) + + assert.equal(info.clients[0].state, 3) + assert.deepEqual(info.connected_clients, []) +} + function assertBooleanFieldValuesPreserved() { const input = allFieldFixture() const normalized = normalizeNetworkConfig(input) @@ -526,6 +575,8 @@ function assertNumberBoundaries() { const tests = [ assertFixtureCoversGeneratedFields, assertFullFieldRoundTrip, + assertLegacyVpnPortalFieldsReachBackendValidation, + assertVpnPortalRpcJsonNormalization, assertBooleanFieldValuesPreserved, assertEnumCompatibility, assertAclDefaultsAndExplicitZero, diff --git a/easytier-web/frontend-lib/src/components/Config.vue b/easytier-web/frontend-lib/src/components/Config.vue index 3d5c9cab..3428a8b1 100644 --- a/easytier-web/frontend-lib/src/components/Config.vue +++ b/easytier-web/frontend-lib/src/components/Config.vue @@ -1,5 +1,6 @@