feat(vpn): multi-client WireGuard portal with attached peers (#2502)

* feat(peer): support protocol-agnostic attached peers

Add locally attached peers backed by independent, peer-level portable
managers and authenticated in-process ring connections. Carry trusted
connection provenance through packet admission so attached relay
privileges cannot be forged through packet headers.

Let every peer manager own ACL loading, sanitized policy updates, route
refresh, and runtime cleanup. In Secure Mode, grant attached identities
ephemeral credentials instead of sharing administrator and group secrets.

* feat(vpn): add reusable attached-peer portal runtime

Add a protocol-neutral portal runtime that converts authenticated client
sessions into attached EasyTier peers. Own per-client generations,
status, packet forwarding, address translation, and peer cleanup without
knowing the transport protocol.

Add transactional IPv4 source and destination rewriting with correct
IPv4, TCP, UDP, ICMP, and quoted-packet checksum updates. Keep the old
production portal path temporarily active until the WireGuard adapter is
migrated in the next change.

* feat(wireguard): attach named clients through peer portal

Replace the monolithic WireGuard portal with a native adapter that owns
key derivation, UDP demultiplexing, reauthentication, roaming, and
bounded per-client packet queues. Hand authenticated sessions to the
generic portal runtime for peer lifecycle and IPv4 translation.

Move portal configuration into the core instance model, require a
dedicated server key, and preserve existing listener, CLI, and runtime
configuration behavior. Reject runtime address conflicts before
publishing shared configuration.

* feat(vpn): expose per-client portal status

Project configured clients and their runtime state through the portal
RPC, including generated client configuration, listener, peer identity,
endpoint, tunnel address, ACL groups, and errors. Keep private client
configuration out of the broad instance-info response and expose the
explicit RPC through the CLI and Tauri bridge.

* feat(vpn): add portal configuration to web clients

Expose WireGuard portal listener, key, client, ACL group, and runtime
status fields in the shared frontend library, Web dashboard, and Tauri
client. Preserve UUID and uint64 values across protobuf JSON
boundaries, keep dynamic client editor rows stable, and document the
portal workflow.

* test(vpn): cover multi-client and roaming WireGuard portals

Add two three-node integration tests for the WireGuard VPN portal.

The multi-client test connects two kernel WireGuard clients from
separate network namespaces, verifies per-client connectivity to mesh
nodes, and exercises cross-client traffic that runs the IPv4 source
and destination translation in both directions. A TCP echo exchange
through the portal additionally covers the TCP pseudo-header checksum
rewrite path that ICMP-only ping tests miss, and portal status
snapshots must report both clients online with distinct peer ids and
correctly learned tunnel addresses.

The roaming test swaps the client namespace address (delete the old
address, then add the new one) so the kernel WireGuard source cache is
invalidated and the client keeps sending under the same session from
the new source, exactly like a real network change. The portal must
update the client endpoint on the same peer id via the data path
(same generation, no re-handshake, no detach/reconnect) while
connectivity to mesh nodes is preserved.

Supporting changes: run_wireguard_client now takes an interface name,
and the shared namespace topology gains net_f (10.1.2.5) on the portal
bridge for the second client.
This commit is contained in:
KKRainbow
2026-08-21 10:59:05 +08:00
committed by GitHub
parent 57eb6908f4
commit 62e4fd15e9
63 changed files with 7759 additions and 1216 deletions
+19
View File
@@ -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
Generated
+2
View File
@@ -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",
+10 -6
View File
@@ -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
+9 -6
View File
@@ -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 地址。
#### 自建公共共享节点
+13 -5
View File
@@ -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()
+118 -26
View File
@@ -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<Vec<PeerConfig>, anyhow::Erro
}
impl NetworkConfigExt for NetworkConfig {
#[allow(deprecated)]
fn gen_config(&self) -> Result<TomlConfigLoader, anyhow::Error> {
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::<Result<Vec<_>, 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());
}
}
+47 -3
View File
@@ -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<Option<Acl>> {
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<Ipv4Cidr>,
pub pinned_peers: Vec<(url::Url, Option<String>)>,
pub peer_group_memberships: Vec<PeerGroupIdentity>,
pub acl_group_declarations: Vec<PeerGroupIdentity>,
@@ -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());
}
}
+9
View File
@@ -145,6 +145,15 @@ impl CoreRuntimeConfigStore {
pub fn subscribe_service_runtime_changes(&self) -> tokio::sync::watch::Receiver<u64> {
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)]
+162 -5
View File
@@ -276,6 +276,9 @@ pub trait ConfigLoader: Send + Sync {
fn set_network_config_source(&self, _source: Option<ConfigSource>) {}
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<String>,
#[serde(default)]
pub clients: Vec<VpnPortalClientConfig>,
}
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(|_| "<redacted>"),
)
.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<String>,
}
#[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 = "<redacted>";
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, anyhow::Error> {
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("<redacted>").count(), 4);
}
#[test]
fn hostname_normalization_is_portable_and_has_no_host_fallback() {
let absent = TomlConfig::default();
@@ -14,14 +14,12 @@ use crate::{
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub(crate) struct ProxyCidrConfigSnapshot {
pub manual_routes: Option<BTreeSet<Ipv4Cidr>>,
pub vpn_portal_cidr: Option<Ipv4Cidr>,
}
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<Ipv4Cidr>,
peer_routes: BTreeSet<Ipv4Cidr>,
config: ProxyCidrConfigSnapshot,
) -> BTreeSet<Ipv4Cidr> {
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"])
);
}
}
+9 -712
View File
@@ -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<V> {
endpoint_addr: Option<url::Url>,
value: V,
}
impl<V> VpnPortalClient<V> {
pub fn endpoint_addr(&self) -> Option<&url::Url> {
self.endpoint_addr.as_ref()
}
pub fn value(&self) -> &V {
&self.value
}
}
pub struct VpnPortalClientTable<V> {
entries: DashMap<Ipv4Addr, Arc<VpnPortalClient<V>>>,
}
impl<V> Default for VpnPortalClientTable<V> {
fn default() -> Self {
Self {
entries: DashMap::new(),
}
}
}
impl<V> VpnPortalClientTable<V> {
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<Option<url::Url>> {
self.entries
.iter()
.map(|entry| entry.value().endpoint_addr.clone())
.collect()
}
pub fn route_peer_packet(&self, packet: &ZCPacket) -> VpnPortalPeerPacketRoute<V> {
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<VpnPortalClient<V>>) {
self.entries.insert(address, client);
}
fn remove_if_current(&self, address: &Ipv4Addr, client: &Arc<VpnPortalClient<V>>) -> 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<V> {
Pass,
Drop,
Deliver {
destination: Ipv4Addr,
client: Arc<VpnPortalClient<V>>,
},
}
#[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<V> {
table: Arc<VpnPortalClientTable<V>>,
client: Arc<VpnPortalClient<V>>,
registered_ip: Option<Ipv4Addr>,
}
pub type VpnPortalListener = Box<dyn SocketListener<Accepted = Box<dyn Tunnel>>>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VpnPortalClientConfigPlan {
pub client_cidr: Ipv4Cidr,
pub allowed_ips: Vec<String>,
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<String>,
}
#[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<Vec<VpnPortalListener>>;
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<url::Url>,
_pipeline: PipelineRegistrationGuard,
}
struct VpnPortalSessionEventGuard {
events: Arc<dyn CoreEventSink>,
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<VpnPortalClientTable<MpscTunnelSender>>,
}
#[async_trait]
impl PeerPacketFilter for VpnPortalPeerPacketFilter {
async fn try_process_packet_from_peer(&self, packet: ZCPacket) -> Option<ZCPacket> {
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<PeerManagerCore>,
runtime_config: CoreRuntimeConfigStore,
host: Option<Arc<dyn VpnPortalHost>>,
events: Arc<dyn CoreEventSink>,
clients: Arc<VpnPortalClientTable<MpscTunnelSender>>,
runtime: Mutex<Option<VpnPortalRuntime>>,
}
impl VpnPortalModule {
pub fn new(
peer_manager: Arc<PeerManagerCore>,
runtime_config: CoreRuntimeConfigStore,
host: Option<Arc<dyn VpnPortalHost>>,
events: Arc<dyn CoreEventSink>,
) -> Arc<Self> {
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<PeerManagerCore>,
clients: Arc<VpnPortalClientTable<MpscTunnelSender>>,
events: Arc<dyn CoreEventSink>,
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<dyn Tunnel>,
peer_manager: Arc<PeerManagerCore>,
clients: Arc<VpnPortalClientTable<MpscTunnelSender>>,
events: Arc<dyn CoreEventSink>,
) {
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<VpnPortalClientConfigPlan> {
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::<Vec<_>>();
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<V> VpnPortalClientSession<V> {
pub fn new(
table: Arc<VpnPortalClientTable<V>>,
endpoint_addr: Option<url::Url>,
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<VpnPortalClientPacket> {
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<Ipv4Addr> {
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<V> Drop for VpnPortalClientSession<V> {
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<u8> {
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
));
}
}
@@ -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<ChecksumField>,
}
#[derive(Clone, Copy)]
enum TransportPlan {
HeaderOnly,
WithPseudoHeaderChecksum(ChecksumField),
Icmp {
checksum_offset: usize,
quoted: Option<QuotedIpv4Plan>,
},
}
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<Ipv4Layout, Ipv4TranslationError> {
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<TransportPlan, Ipv4TranslationError> {
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<QuotedIpv4Plan, Ipv4TranslationError> {
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<u8> {
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<u8> {
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<u8> {
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<u8> {
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, &quoted);
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, &quoted);
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,
})
);
}
}
File diff suppressed because it is too large Load Diff
@@ -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(
+15 -4
View File
@@ -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,
+23 -76
View File
@@ -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<PortalRuntimeConfig>,
}
#[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<Option<crate::proto::acl::Acl>> {
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<PeerRuntimeSnapshot>,
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<Arc<dyn PublicIpv6ProviderPlatform>>,
#[cfg(feature = "vpn-portal")]
pub vpn_portal: Option<Arc<dyn VpnPortalHost>>,
pub vpn_portal: Option<Arc<dyn PortalHost>>,
}
impl<H> CoreHostAdapters<H>
@@ -441,7 +408,7 @@ where
#[cfg(feature = "public-ipv6-provider")]
public_ipv6_provider: PublicIpv6ProviderRuntime,
#[cfg(feature = "vpn-portal")]
vpn_portal: Arc<VpnPortalModule>,
vpn_portal: Arc<PortalModule>,
#[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<H>,
) -> anyhow::Result<Arc<Self>> {
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) {
+154 -38
View File
@@ -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<Arc<CoreInstance<TestHost>>> {
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::<cidr::Ipv4Inet>().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::<IpAddr>().unwrap()
);
instance.peer_manager.clear_resources().await;
}
#[cfg(feature = "web-client")]
#[tokio::test]
@@ -1,5 +1,5 @@
use crate::{
gateway::vpn_portal::VpnPortalInfoSnapshot,
gateway::vpn_portal::PortalInfoSnapshot,
instance::{CoreInstance, CoreInstanceHost},
};
@@ -7,7 +7,7 @@ impl<H> CoreInstance<H>
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
}
}
@@ -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<H>(
instance: &CoreInstance<H>,
) -> anyhow::Result<NetworkInstanceRunningInfo>
@@ -50,10 +51,6 @@ where
.map(Into::into)
.collect::<Vec<_>>();
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(),
@@ -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<F, H> VpnPortalRpc for InstanceManagementRpc<F>
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"));
}
}
File diff suppressed because it is too large Load Diff
+7
View File
@@ -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<uuid::Uuid> {
self.conns
.iter()
+28 -5
View File
@@ -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<PeerId>,
@@ -319,21 +320,30 @@ impl Debug for PeerConn {
}
impl PeerConn {
#[cfg(test)]
pub(crate) fn new(
my_peer_id: PeerId,
context: ArcPeerContext,
tunnel: Box<dyn Tunnel>,
peer_session_store: Arc<PeerSessionStore>,
) -> 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<dyn Tunnel>,
peer_id_hint: Option<PeerId>,
peer_session_store: Arc<PeerSessionStore>,
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;
}
+5
View File
@@ -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
}
+4 -16
View File
@@ -76,7 +76,6 @@ pub struct PeerRuntimeSnapshotInput {
pub host_routing: HostRoutingPolicy,
pub acl: Option<Acl>,
pub easytier_version: String,
pub vpn_portal_cidr: Option<Ipv4Cidr>,
pub pinned_peers: Vec<(url::Url, Option<String>)>,
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<StatsManager> {
self.stats_manager.clone()
}
@@ -604,10 +605,6 @@ pub(crate) trait PeerContext: Send + Sync {
Vec::new()
}
fn vpn_portal_cidr(&self) -> Option<Ipv4Cidr> {
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<Ipv4Cidr> {
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![(
+130 -9
View File
@@ -116,6 +116,7 @@ pub trait CredentialStorage: Send + Sync + 'static {
pub(crate) struct CredentialManager {
credentials: Mutex<HashMap<String, CredentialEntry>>,
ephemeral_credentials: Mutex<HashMap<uuid::Uuid, CredentialEntry>>,
storage: Option<Arc<dyn CredentialStorage>>,
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<String>,
allow_relay: bool,
allowed_proxy_cidrs: Vec<String>,
reusable: bool,
) -> Result<uuid::Uuid, String> {
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<String>,
) -> Option<bool> {
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<bool, String> {
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<TrustedCredentialPubkeyProof> {
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::<Vec<_>>();
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<CredentialInfo> {
@@ -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());
}
}
@@ -93,6 +93,10 @@ impl ForeignNetworkClient {
})));
}
pub(crate) fn stop(&self) {
self.task.lock().unwrap().take();
}
pub fn get_peer_map(&self) -> Arc<PeerMap> {
self.peer_map.clone()
}
+115 -15
View File
@@ -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<ZCPacket>;
pub type PacketRecvChanReceiver = tokio::sync::mpsc::Receiver<ZCPacket>;
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<PeerPacketEnvelope>);
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<ZCPacket>> {
self.0
.send(PeerPacketEnvelope::local(packet))
.await
.map_err(|error| SendError(error.0.packet))
}
}
pub struct PacketRecvChanReceiver(tokio::sync::mpsc::Receiver<PeerPacketEnvelope>);
impl PacketRecvChanReceiver {
pub async fn recv(&mut self) -> Option<ZCPacket> {
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<ZCPacket>> {
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<PeerPacketEnvelope, anyhow::Error> {
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<ZCPacket, anyhow::Error> {
match packet_recv_chan_receiver.try_recv() {
Ok(packet) => Ok(packet),
Err(TryRecvError::Empty) => packet_recv_chan_receiver
.recv()
recv_packet_envelope_from_chan(packet_recv_chan_receiver)
.await
.ok_or(anyhow::anyhow!("recv_packet_from_chan failed")),
Err(TryRecvError::Disconnected) => Err(anyhow::anyhow!("recv_packet_from_chan failed")),
}
.map(|envelope| envelope.packet)
}
#[async_trait::async_trait]
+521 -25
View File
@@ -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<PeerGroupIdentity> {
let configured = configured_groups.iter().collect::<BTreeSet<_>>();
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::<BTreeSet<_>>();
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<RwLock<Vec<IpAddr>>>,
acl_filter: Arc<AclFilter>,
context: Arc<CorePeerContext>,
runtime_config: CoreRuntimeConfigStore,
runtime_config_update: Mutex<()>,
is_secure_mode_enabled: bool,
route: ArcRoute,
traffic_metrics: Arc<TrafficMetricRecorder>,
@@ -771,6 +819,7 @@ impl PeerManagerCore {
credential_storage: Option<Arc<dyn CredentialStorage>>,
foreign_rpc_registrar: Arc<dyn ForeignNetworkRpcRegistrar>,
) -> anyhow::Result<Self> {
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<String>,
allow_relay: bool,
allowed_proxy_cidrs: Vec<String>,
reusable: bool,
) -> anyhow::Result<uuid::Uuid> {
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<String>,
) -> Option<bool> {
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<StatsManager> {
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<Arc<CoreInstanceRuntimeConfig>> {
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<Self>,
source: CoreRuntimeConfigStore,
configured_groups: Vec<String>,
) -> 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<PeerSessionStore> {
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<dyn Tunnel>,
) -> 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<dyn Tunnel>,
@@ -1363,6 +1598,15 @@ impl PeerManagerCore {
.await
}
pub(crate) async fn add_attached_ring_tunnel_as_server(
&self,
tunnel: Box<dyn Tunnel>,
) -> 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<PeerId>,
) -> 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<dyn Tunnel>,
is_directly_connected: bool,
peer_id_hint: Option<PeerId>,
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<dyn Tunnel>,
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<dyn Tunnel>,
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<std::net::Ipv4Addr>,
@@ -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()),
+80
View File
@@ -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());
+30
View File
@@ -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<Option<VpnPortalInfo>, String> {
let instance_id = instance_id
.parse::<uuid::Uuid>()
.map_err(|e| e.to_string())?;
let client_manager = get_client_manager!()?;
let client = client_manager
.rpc_manager
.rpc_client()
.scoped_client::<VpnPortalRpcClientFactory<BaseController>>(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,
+2
View File
@@ -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<typeof import('./composables/backend')['getEasytierVersion']>
readonly getNetworkMetas: UnwrapRef<typeof import('./composables/backend')['getNetworkMetas']>
readonly getServiceStatus: UnwrapRef<typeof import('./composables/backend')['getServiceStatus']>
readonly getVpnPortalInfo: UnwrapRef<typeof import('./composables/backend')['getVpnPortalInfo']>
readonly h: UnwrapRef<typeof import('vue')['h']>
readonly initMobileVpnService: UnwrapRef<typeof import('./composables/mobile_vpn')['initMobileVpnService']>
readonly initRpcConnection: UnwrapRef<typeof import('./composables/backend')['initRpcConnection']>
+5
View File
@@ -66,6 +66,11 @@ export async function collectNetworkInfo(instanceId: string) {
return await invoke<Api.CollectNetworkInfoResponse>('collect_network_info', { instanceId })
}
export async function getVpnPortalInfo(instanceId: string) {
const info = await invoke<NetworkTypes.VpnPortalInfo | undefined>('get_vpn_portal_info', { instanceId })
return info ? NetworkTypes.normalizeVpnPortalInfo(info) : undefined
}
export async function setLoggingLevel(level: string) {
return await invoke('set_logging_level', { level })
}
+3
View File
@@ -11,6 +11,9 @@ export class GUIRemoteClient implements Api.RemoteClient {
async get_network_info(inst_id: string): Promise<NetworkTypes.NetworkInstanceRunningInfo | undefined> {
return backend.collectNetworkInfo(inst_id).then(infos => infos.info?.map?.[inst_id]);
}
async get_vpn_portal_info(inst_id: string): Promise<NetworkTypes.VpnPortalInfo | undefined> {
return backend.getVpnPortalInfo(inst_id);
}
async list_network_instance_ids(): Promise<Api.ListNetworkInstanceIdResponse> {
return backend.listNetworkInstanceIds();
}
+6 -1
View File
@@ -124,7 +124,12 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
.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/"])?;
+24 -2
View File
@@ -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; }
+18 -5
View File
@@ -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;
}
+27
View File
@@ -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(|_| "<redacted>"),
)
.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("<redacted>"));
assert!(!debug.contains("private-key-material"));
}
#[derive(Clone, Default)]
struct WebClientServiceJsonCallHandler;
@@ -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,
@@ -1,5 +1,6 @@
<script setup lang="ts">
import { AutoComplete, Button, Checkbox, Dialog, Divider, InputNumber, InputText, Panel, Password, SelectButton, ToggleButton } from 'primevue'
import { v4 as uuidv4 } from 'uuid'
import { AutoComplete, Button, Checkbox, Dialog, Divider, InputNumber, InputText, MultiSelect, Panel, Password, SelectButton, ToggleButton } from 'primevue'
import InputGroup from 'primevue/inputgroup'
import InputGroupAddon from 'primevue/inputgroupaddon'
import {
@@ -7,7 +8,9 @@ import {
DEFAULT_NETWORK_CONFIG,
NetworkConfig,
normalizeNetworkConfig,
removeRow
removeRow,
type VpnPortalClientConfig,
type VpnPortalConfig,
} from '../types/network'
import { computed, ref, onMounted, onUnmounted, watch } from 'vue'
import { useI18n } from 'vue-i18n'
@@ -196,6 +199,52 @@ const instanceRecvBpsLimitInput = computed<string>({
}
},
})
function defaultVpnPortalConfig(): VpnPortalConfig {
return {
wireguard_listen: '0.0.0.0:22022',
clients: [],
}
}
const vpnPortalEnabled = computed({
get: () => curNetwork.value.vpn_portal_config !== undefined,
set: (enabled: boolean) => {
curNetwork.value.vpn_portal_config = enabled ? defaultVpnPortalConfig() : undefined
},
})
const vpnPortalConfig = computed(() => curNetwork.value.vpn_portal_config ?? defaultVpnPortalConfig())
const vpnPortalPrivateKey = computed({
get: () => vpnPortalConfig.value.wireguard_private_key ?? '',
set: (value: string | null | undefined) => {
vpnPortalConfig.value.wireguard_private_key = value && value.length > 0 ? value : undefined
},
})
const vpnPortalGroupOptions = computed(() => (
curNetwork.value.acl?.acl_v1?.group?.declares ?? []
).map((group) => group.group_name))
const vpnPortalClientViewKeys = new WeakMap<VpnPortalClientConfig, string>()
function vpnPortalClientViewKey(client: VpnPortalClientConfig): string {
let key = vpnPortalClientViewKeys.get(client)
if (!key) {
key = uuidv4()
vpnPortalClientViewKeys.set(client, key)
}
return key
}
function addVpnPortalClient() {
vpnPortalConfig.value.clients.push({ name: '', virtual_ip: '', groups: [] })
}
function removeVpnPortalClient(index: number) {
vpnPortalConfig.value.clients.splice(index, 1)
}
</script>
<template>
@@ -295,25 +344,57 @@ const instanceRecvBpsLimitInput = computed<string>({
<div class="flex flex-row gap-x-9 flex-wrap ">
<div class="flex flex-col gap-2 grow">
<label for="username">VPN Portal</label>
<ToggleButton v-model="curNetwork.enable_vpn_portal" on-icon="pi pi-check" off-icon="pi pi-times"
<label>VPN Portal</label>
<ToggleButton v-model="vpnPortalEnabled" on-icon="pi pi-check" off-icon="pi pi-times"
:on-label="t('off_text')" :off-label="t('on_text')" class="w-48" />
<div v-if="curNetwork.enable_vpn_portal" class="items-center flex flex-row gap-x-4">
<div class="flex flex-row gap-x-9 flex-wrap w-full">
<div class="flex flex-col gap-2 basis-8/12 grow">
<InputGroup>
<InputText v-model="curNetwork.vpn_portal_client_network_addr"
:placeholder="t('vpn_portal_client_network')" />
<InputGroupAddon>
<span>/{{ curNetwork.vpn_portal_client_network_len }}</span>
</InputGroupAddon>
</InputGroup>
<div v-if="vpnPortalEnabled" class="flex flex-col gap-3 w-full">
<div class="flex flex-row gap-x-9 gap-y-3 flex-wrap w-full">
<div class="flex flex-col gap-2 basis-5/12 grow">
<label for="vpn_portal_wireguard_listen">{{ t('vpn_portal_wireguard_listen') }}</label>
<InputText id="vpn_portal_wireguard_listen" v-model="vpnPortalConfig.wireguard_listen"
:placeholder="t('vpn_portal_wireguard_listen_placeholder')" />
</div>
<div class="flex flex-col gap-2 basis-3/12 grow">
<InputNumber v-model="curNetwork.vpn_portal_listen_port" :allow-empty="false" :format="false"
:min="0" :max="65535" fluid />
<div class="flex flex-col gap-2 basis-5/12 grow">
<label for="vpn_portal_wireguard_private_key">{{ t('vpn_portal_wireguard_private_key') }}</label>
<Password id="vpn_portal_wireguard_private_key"
v-model="vpnPortalPrivateKey"
:placeholder="t('vpn_portal_wireguard_private_key_placeholder')"
toggleMask :feedback="false" fluid />
</div>
</div>
<div class="flex items-center justify-between gap-3">
<label>{{ t('vpn_portal_clients') }}</label>
<Button icon="pi pi-plus" :label="t('vpn_portal_add_client')" severity="secondary" size="small"
:disabled="vpnPortalConfig.clients.length >= 64"
@click="addVpnPortalClient" />
</div>
<div v-if="vpnPortalConfig.clients.length === 0"
class="text-sm text-surface-500 dark:text-surface-400">
{{ t('vpn_portal_no_clients') }}
</div>
<div v-for="(client, index) in vpnPortalConfig.clients" :key="vpnPortalClientViewKey(client)"
class="flex flex-row gap-3 flex-wrap items-end rounded border border-surface-200 dark:border-surface-700 p-3">
<div class="flex flex-col gap-2 grow basis-3/12">
<label :for="`vpn_portal_client_name_${index}`">{{ t('vpn_portal_client_name') }}</label>
<InputText :id="`vpn_portal_client_name_${index}`" v-model="client.name"
:placeholder="t('vpn_portal_client_name_placeholder')" />
</div>
<div class="flex flex-col gap-2 grow basis-3/12">
<label :for="`vpn_portal_client_virtual_ip_${index}`">{{ t('vpn_portal_client_virtual_ip') }}</label>
<InputText :id="`vpn_portal_client_virtual_ip_${index}`" v-model="client.virtual_ip"
:placeholder="t('vpn_portal_client_virtual_ip_placeholder')" />
</div>
<div class="flex flex-col gap-2 grow basis-4/12">
<label :for="`vpn_portal_client_groups_${index}`">{{ t('vpn_portal_client_groups') }}</label>
<MultiSelect :input-id="`vpn_portal_client_groups_${index}`" v-model="client.groups"
:options="vpnPortalGroupOptions" appendTo="self" filter fluid
:placeholder="t('vpn_portal_client_groups_placeholder')" />
</div>
<Button icon="pi pi-trash" severity="danger" text rounded
:aria-label="t('vpn_portal_remove_client')" @click="removeVpnPortalClient(index)" />
</div>
</div>
</div>
</div>
@@ -588,6 +588,7 @@ onUnmounted(() => {
</div>
<Status v-if="curNetworkInfo && curNetworkInfo.error_msg === ''" v-bind:cur-network-inst="curNetworkInfo"
:api="api"
class="mb-4">
</Status>
<Message v-else-if="curNetworkInfo?.error_msg" severity="error" class="mb-4">{{
@@ -1,6 +1,7 @@
<script setup lang="ts">
import { useTimeAgo } from '@vueuse/core'
import { NetworkInstance, type TunnelInfo, type NodeInfo, type PeerRoutePair } from '../types/network'
import { NetworkInstance, VpnPortalClientState, type TunnelInfo, type NodeInfo, type PeerRoutePair, type VpnPortalClientInfo, type VpnPortalInfo } from '../types/network'
import type { RemoteClient } from '../modules/api'
import { useI18n } from 'vue-i18n';
import { computed, onMounted, onUnmounted, ref } from 'vue';
import { ipv4InetToString, ipv4ToString, ipv6ToString } from '../modules/utils';
@@ -10,6 +11,7 @@ import NetworkChart from './NetworkChart.vue';
const props = defineProps<{
curNetworkInst: NetworkInstance | null,
api: RemoteClient,
}>()
const { t } = useI18n()
@@ -327,16 +329,66 @@ onUnmounted(() => {
const dialogVisible = ref(false)
const dialogContent = ref<any>('')
const dialogHeader = ref('event_log')
const vpnPortalInfo = ref<VpnPortalInfo>()
const vpnPortalClients = computed(() => vpnPortalInfo.value?.clients ?? [])
const vpnPortalLoading = ref(false)
const vpnPortalError = ref('')
const copiedVpnPortalClient = ref('')
function showVpnPortalConfig() {
const my_node_info = myNodeInfo.value
if (!my_node_info)
async function showVpnPortalConfig() {
const instanceId = props.curNetworkInst?.instance_id
if (!instanceId)
return
const url = 'https://www.wireguardconfig.com/qrcode'
dialogContent.value = `${my_node_info.vpn_portal_cfg}\n\n # can generate QR code: ${url}`
dialogHeader.value = 'vpn_portal_config'
dialogVisible.value = true
vpnPortalInfo.value = undefined
vpnPortalError.value = ''
copiedVpnPortalClient.value = ''
vpnPortalLoading.value = true
try {
vpnPortalInfo.value = await props.api.get_vpn_portal_info(instanceId)
} catch (error) {
console.error('Failed to load VPN Portal information', error)
vpnPortalError.value = t('vpn_portal_load_failed')
} finally {
vpnPortalLoading.value = false
}
}
function vpnPortalStateKey(state: VpnPortalClientState | string): string {
const normalized = typeof state === 'string'
? state.toLowerCase().replace('vpn_portal_client_state_', '')
: VpnPortalClientState[state]?.toLowerCase()
return `vpn_portal_state_${normalized ?? 'unspecified'}`
}
function vpnPortalStateSeverity(state: VpnPortalClientState | string): 'success' | 'warn' | 'danger' | 'secondary' {
const key = vpnPortalStateKey(state)
if (key.endsWith('online')) return 'success'
if (key.endsWith('connecting')) return 'warn'
if (key.endsWith('error')) return 'danger'
return 'secondary'
}
async function copyVpnPortalClientConfig(client: VpnPortalClientInfo) {
try {
if (navigator.clipboard?.writeText) {
await navigator.clipboard.writeText(client.client_config)
} else {
const textarea = document.createElement('textarea')
textarea.value = client.client_config
textarea.style.position = 'fixed'
textarea.style.opacity = '0'
document.body.appendChild(textarea)
textarea.select()
document.execCommand('copy')
textarea.remove()
}
copiedVpnPortalClient.value = client.name
} catch (error) {
console.error('Failed to copy VPN Portal client config', error)
}
}
function showEventLogs() {
@@ -354,8 +406,52 @@ function showEventLogs() {
<div class="frontend-lib">
<Dialog v-model:visible="dialogVisible" modal :header="t(dialogHeader)" class="w-full h-auto max-h-full"
:baseZIndex="2000">
<ScrollPanel v-if="dialogHeader === 'vpn_portal_config'">
<pre>{{ dialogContent }}</pre>
<ScrollPanel v-if="dialogHeader === 'vpn_portal_config'" class="max-h-[75vh] pr-3">
<div v-if="vpnPortalLoading" class="py-8 text-center text-surface-500">
{{ t('web.device_management.loading_network_status') }}
</div>
<div v-else-if="vpnPortalError" class="py-4 text-red-500">
{{ vpnPortalError }}
</div>
<div v-else-if="!vpnPortalInfo || ((!vpnPortalInfo.vpn_type || vpnPortalInfo.vpn_type === 'null') && vpnPortalClients.length === 0)"
class="py-4 text-surface-500">
{{ t('vpn_portal_not_configured') }}
</div>
<div v-else class="flex flex-col gap-4">
<div class="flex flex-wrap gap-x-6 gap-y-2 text-sm">
<span v-if="vpnPortalInfo.vpn_type"><strong>{{ t('vpn_portal_type') }}:</strong>
{{ vpnPortalInfo.vpn_type }}</span>
<span v-if="vpnPortalInfo.listener"><strong>{{ t('vpn_portal_listener') }}:</strong>
{{ vpnPortalInfo.listener }}</span>
</div>
<div v-for="client in vpnPortalClients" :key="client.name"
class="rounded border border-surface-200 dark:border-surface-700 p-4">
<div class="mb-3 flex flex-wrap items-center justify-between gap-2">
<div class="font-semibold">{{ client.name }} · {{ client.virtual_ip }}</div>
<Tag :severity="vpnPortalStateSeverity(client.state)"
:value="t(vpnPortalStateKey(client.state))" />
</div>
<div class="mb-3 grid gap-x-6 gap-y-1 text-sm sm:grid-cols-2">
<span v-if="client.groups.length"><strong>{{ t('vpn_portal_client_groups') }}:</strong>
{{ client.groups.join(', ') }}</span>
<span v-if="client.peer_id !== undefined"><strong>{{ t('vpn_portal_peer_id') }}:</strong>
{{ client.peer_id }}</span>
<span v-if="client.endpoint"><strong>{{ t('vpn_portal_endpoint') }}:</strong>
{{ client.endpoint }}</span>
<span v-if="client.tunnel_ip"><strong>{{ t('vpn_portal_tunnel_ip') }}:</strong>
{{ client.tunnel_ip }}</span>
<span v-if="client.error" class="text-red-500 sm:col-span-2">{{ client.error }}</span>
</div>
<div class="mb-2 flex items-center justify-between gap-3">
<label class="font-medium">{{ t('vpn_portal_client_config') }}</label>
<Button size="small" severity="secondary" icon="pi pi-copy"
:label="copiedVpnPortalClient === client.name ? t('config_copied') : t('vpn_portal_copy_client_config')"
@click="copyVpnPortalClientConfig(client)" />
</div>
<pre class="max-w-full overflow-x-auto whitespace-pre-wrap break-all rounded bg-surface-100 p-3 text-xs dark:bg-surface-800">{{ client.client_config }}</pre>
</div>
</div>
</ScrollPanel>
<Timeline v-else :value="dialogContent">
<template #opposite="slotProps">
+29 -3
View File
@@ -18,9 +18,20 @@ network_secret: 网络密码
public_server_url: 公共服务器地址
peer_urls: 对等节点地址
proxy_cidrs: 子网代理CIDR
enable_vpn_portal: 启用VPN门户
vpn_portal_listen_port: 监听端口
vpn_portal_client_network: 客户端子网
vpn_portal_wireguard_listen: WireGuard 监听地址
vpn_portal_wireguard_listen_placeholder: 例如:0.0.0.0:22022
vpn_portal_wireguard_private_key: WireGuard 服务端私钥
vpn_portal_wireguard_private_key_placeholder: 必填 Base64 密钥(可用 wg genkey 生成)
vpn_portal_clients: WireGuard 客户端
vpn_portal_add_client: 添加客户端
vpn_portal_no_clients: 尚未配置客户端
vpn_portal_client_name: 客户端名称
vpn_portal_client_name_placeholder: 例如:alice-phone
vpn_portal_client_virtual_ip: 虚拟网地址
vpn_portal_client_virtual_ip_placeholder: 例如:10.126.126.10
vpn_portal_client_groups: ACL 组
vpn_portal_client_groups_placeholder: 选择 ACL 组
vpn_portal_remove_client: 删除客户端
dev_name: TUN接口名称
advanced_settings: 高级设置
basic_settings: 基础设置
@@ -87,6 +98,21 @@ upload: 上传
download: 下载
show_vpn_portal_config: 显示VPN门户配置
vpn_portal_config: VPN门户配置
vpn_portal_not_configured: 当前节点未配置 VPN 门户
vpn_portal_load_failed: VPN 门户信息加载失败
vpn_portal_listener: 监听地址
vpn_portal_type: 类型
vpn_portal_state: 状态
vpn_portal_peer_id: 节点 ID
vpn_portal_endpoint: 客户端端点
vpn_portal_tunnel_ip: 客户端隧道地址
vpn_portal_client_config: 客户端配置
vpn_portal_copy_client_config: 复制客户端配置
vpn_portal_state_unspecified: 未知
vpn_portal_state_offline: 离线
vpn_portal_state_connecting: 连接中
vpn_portal_state_online: 在线
vpn_portal_state_error: 错误
show_event_log: 显示事件日志
event_log: 事件日志
peer_info: 节点信息
+29 -3
View File
@@ -18,9 +18,20 @@ network_secret: Network Secret
public_server_url: Public Server URL
peer_urls: Peer URLs
proxy_cidrs: Subnet Proxy CIDRs
enable_vpn_portal: Enable VPN Portal
vpn_portal_listen_port: VPN Portal Listen Port
vpn_portal_client_network: Client Sub Network
vpn_portal_wireguard_listen: WireGuard Listen Address
vpn_portal_wireguard_listen_placeholder: "Example: 0.0.0.0:22022"
vpn_portal_wireguard_private_key: WireGuard Server Private Key
vpn_portal_wireguard_private_key_placeholder: Required base64 key (generate with wg genkey)
vpn_portal_clients: WireGuard Clients
vpn_portal_add_client: Add Client
vpn_portal_no_clients: No clients configured
vpn_portal_client_name: Client Name
vpn_portal_client_name_placeholder: "Example: alice-phone"
vpn_portal_client_virtual_ip: Virtual Network Address
vpn_portal_client_virtual_ip_placeholder: "Example: 10.126.126.10"
vpn_portal_client_groups: ACL Groups
vpn_portal_client_groups_placeholder: Select ACL groups
vpn_portal_remove_client: Remove Client
dev_name: TUN interface name
advanced_settings: Advanced Settings
basic_settings: Basic Settings
@@ -86,6 +97,21 @@ upload: Upload
download: Download
show_vpn_portal_config: Show VPN Portal Config
vpn_portal_config: VPN Portal Config
vpn_portal_not_configured: VPN Portal is not configured on this node
vpn_portal_load_failed: Failed to load VPN Portal information
vpn_portal_listener: Listener
vpn_portal_type: Type
vpn_portal_state: State
vpn_portal_peer_id: Peer ID
vpn_portal_endpoint: Client Endpoint
vpn_portal_tunnel_ip: Client Tunnel Address
vpn_portal_client_config: Client Config
vpn_portal_copy_client_config: Copy Client Config
vpn_portal_state_unspecified: Unknown
vpn_portal_state_offline: Offline
vpn_portal_state_connecting: Connecting
vpn_portal_state_online: Online
vpn_portal_state_error: Error
show_event_log: Show Event Log
event_log: Event Log
peer_info: Peer Info
+2 -1
View File
@@ -1,5 +1,5 @@
import { UUID } from './utils';
import { NetworkConfig, NetworkInstanceRunningInfo } from '../types/network';
import { NetworkConfig, NetworkInstanceRunningInfo, VpnPortalInfo } from '../types/network';
export interface ValidateConfigResponse {
toml_config: string;
@@ -57,6 +57,7 @@ export interface RemoteClient {
validate_config(config: NetworkConfig): Promise<ValidateConfigResponse>;
run_network(config: NetworkConfig, save: boolean): Promise<undefined>;
get_network_info(inst_id: string): Promise<NetworkInstanceRunningInfo | undefined>;
get_vpn_portal_info(inst_id: string): Promise<VpnPortalInfo | undefined>;
list_network_instance_ids(): Promise<ListNetworkInstanceIdResponse>;
delete_network(inst_id: string): Promise<undefined>;
update_network_instance_state(inst_id: string, disabled: boolean): Promise<undefined>;
@@ -62,6 +62,20 @@ export function UuidToStr(uuid: UUID | null | undefined): string {
return uint32ToUuid(uuid.part1 ?? 0, uuid.part2 ?? 0, uuid.part3 ?? 0, uuid.part4 ?? 0);
}
export function StrToUuid(uuid: string): UUID {
const hex = uuid.replace(/-/g, '');
if (!/^[0-9a-fA-F]{32}$/.test(hex)) {
throw new Error(`Invalid UUID: ${uuid}`);
}
return {
part1: Number.parseInt(hex.slice(0, 8), 16),
part2: Number.parseInt(hex.slice(8, 16), 16),
part3: Number.parseInt(hex.slice(16, 24), 16),
part4: Number.parseInt(hex.slice(24, 32), 16),
};
}
export interface Location {
country: string | undefined;
city: string | undefined;
+24 -8
View File
@@ -5,7 +5,15 @@ import {
type NetworkPeerConfig,
type NetworkConfig as ProtoNetworkConfig,
type PortForwardConfig,
type VpnPortalClientConfig,
type VpnPortalConfig,
} from '../generated/proto/api_manage'
import {
VpnPortalClientState,
VpnPortalInfo as VpnPortalInfoPb,
type VpnPortalClientInfo,
type VpnPortalInfo,
} from '../generated/proto/api_instance'
import {
Action as AclAction,
ChainType as AclChainType,
@@ -26,11 +34,15 @@ import {
import { prepareNetworkConfigForProtoJson } from './networkCompat'
export { AclAction, AclChainType, AclProtocol, CompressionAlgoPb, NatType, NetworkingMethod }
export type { Acl, AclChain, AclRule, AclV1, GroupIdentity, GroupInfo, NetworkPeerConfig, PeerFeatureFlag, PortForwardConfig, SecureModeConfig }
export { VpnPortalClientState }
export type { Acl, AclChain, AclRule, AclV1, GroupIdentity, GroupInfo, NetworkPeerConfig, PeerFeatureFlag, PortForwardConfig, SecureModeConfig, VpnPortalClientConfig, VpnPortalClientInfo, VpnPortalConfig, VpnPortalInfo }
export type NetworkConfig = Omit<
ProtoNetworkConfig,
'instance_id' | 'instance_recv_bps_limit' | 'mtu' | 'networking_method'
| 'instance_id'
| 'instance_recv_bps_limit'
| 'mtu'
| 'networking_method'
> & {
instance_id: string
mtu: number | null
@@ -90,11 +102,6 @@ export function DEFAULT_NETWORK_CONFIG(): NetworkConfig {
proxy_cidrs: [],
enable_vpn_portal: false,
vpn_portal_listen_port: 22022,
vpn_portal_client_network_addr: '',
vpn_portal_client_network_len: 24,
advanced_settings: false,
listener_urls: [
@@ -308,6 +315,12 @@ export function normalizeNetworkConfig(config: NetworkConfig): NetworkConfig {
normalized.exit_nodes ??= []
normalized.mapped_listeners ??= []
normalized.port_forwards ??= []
if (normalized.vpn_portal_config) {
normalized.vpn_portal_config.clients ??= []
normalized.vpn_portal_config.clients.forEach((client) => {
client.groups ??= []
})
}
normalized.acl = config.acl === undefined ? undefined : normalizeAcl(normalized.acl)
return normalized
@@ -330,6 +343,10 @@ export function toBackendNetworkConfig(config: NetworkConfig): NetworkConfig {
}) as unknown as NetworkConfig
}
export function normalizeVpnPortalInfo(info: unknown): VpnPortalInfo {
return VpnPortalInfoPb.fromJson(info as any, { ignoreUnknownFields: true })
}
export interface NetworkInstance {
instance_id: string
@@ -394,7 +411,6 @@ export interface NodeInfo {
}
stun_info: StunInfo
listeners: Url[]
vpn_portal_cfg?: string
peer_id: number
}
@@ -11,7 +11,8 @@ type JsonRecord = Record<string, unknown>
export function prepareNetworkConfigForProtoJson(config: NetworkConfig): NetworkConfig {
const prepared = dropUnsupportedJsonValues(applyLegacyAclDefaults(config)) as NetworkConfig
normalizeLegacyOptionalUint64(prepared as JsonRecord, 'instance_recv_bps_limit')
const record = prepared as JsonRecord
normalizeLegacyOptionalUint64(record, 'instance_recv_bps_limit')
return prepared
}
+111 -10
View File
@@ -43,7 +43,6 @@ const CONFIG_CHECKBOX_FIELDS = [
] as const satisfies readonly (readonly [keyof NetworkConfig, string])[]
const CONFIG_TOGGLE_FIELDS = [
'enable_vpn_portal',
'enable_relay_network_whitelist',
'enable_manual_routes',
'enable_socks5',
@@ -212,6 +211,26 @@ const AutoCompleteStub = defineComponent({
},
})
const MultiSelectStub = defineComponent({
name: 'MultiSelect',
props: {
modelValue: Array,
inputId: String,
appendTo: String,
},
emits: ['update:modelValue'],
setup(props, { attrs, emit }) {
return () => h('input', {
...attrs,
id: props.inputId,
'data-append-to': props.appendTo,
value: (props.modelValue ?? []).join(','),
'data-stub': 'multi-select',
onInput: (event: Event) => emit('update:modelValue', splitList((event.target as HTMLInputElement).value)),
})
},
})
const UrlListInputStub = defineComponent({
name: 'UrlListInput',
props: {
@@ -307,9 +326,15 @@ function makeConfig(): NetworkConfig {
no_tun: true,
hostname: 'host-a',
proxy_cidrs: ['10.10.0.0/16', '172.16.1.0/24'],
enable_vpn_portal: true,
vpn_portal_client_network_addr: '10.144.0.0',
vpn_portal_listen_port: 22023,
vpn_portal_config: {
wireguard_listen: '0.0.0.0:22023',
wireguard_private_key: 'portal-private-key',
clients: [{
name: 'phone-a',
virtual_ip: '10.1.2.10',
groups: ['ops'],
}],
},
listener_urls: ['tcp://0.0.0.0:12010'],
dev_name: 'tun-test',
mtu: 1280,
@@ -354,6 +379,7 @@ function mountConfig(config: NetworkConfig = makeConfig()) {
InputGroupAddon: PassThrough,
InputNumber: InputNumberStub,
InputText: InputTextStub,
MultiSelect: MultiSelectStub,
Panel: PanelStub,
Password: PasswordStub,
SelectButton: SelectButtonStub,
@@ -392,7 +418,11 @@ describe('Config.vue network config projection', () => {
expect(input(wrapper, '#hostname').value).toBe('host-a')
expect(input(wrapper, '#subnet-proxy').value).toBe('10.10.0.0/16,172.16.1.0/24')
expect(input(wrapper, 'input[placeholder="vpn_portal_client_network"]').value).toBe('10.144.0.0')
expect(input(wrapper, '#vpn_portal_wireguard_listen').value).toBe('0.0.0.0:22023')
expect(input(wrapper, '#vpn_portal_wireguard_private_key').value).toBe('portal-private-key')
expect(input(wrapper, '#vpn_portal_client_name_0').value).toBe('phone-a')
expect(input(wrapper, '#vpn_portal_client_virtual_ip_0').value).toBe('10.1.2.10')
expect(input(wrapper, '#vpn_portal_client_groups_0').value).toBe('ops')
expect(input(wrapper, '#dev_name').value).toBe('tun-test')
expect(input(wrapper, '#mtu').value).toBe('1280')
expect(input(wrapper, '#instance_recv_bps_limit').value).toBe('9007199254740993')
@@ -422,7 +452,11 @@ describe('Config.vue network config projection', () => {
await wrapper.find('#disable_ipv6').setValue(false)
await setInput(wrapper, '#hostname', 'host-edited')
await setInput(wrapper, '#subnet-proxy', '10.7.0.0/16,172.17.0.0/16')
await setInput(wrapper, 'input[placeholder="vpn_portal_client_network"]', '10.200.0.0')
await setInput(wrapper, '#vpn_portal_wireguard_listen', '[::]:23000')
await setInput(wrapper, '#vpn_portal_wireguard_private_key', 'edited-private-key')
await setInput(wrapper, '#vpn_portal_client_name_0', 'laptop-a')
await setInput(wrapper, '#vpn_portal_client_virtual_ip_0', '10.1.2.20')
await setInput(wrapper, '#vpn_portal_client_groups_0', 'ops,admin')
await setInput(wrapper, 'input[data-add-label="add_listener_url"]', 'tcp://0.0.0.0:13010')
await setInput(wrapper, '#dev_name', 'tun-edited')
await setInput(wrapper, '#mtu', '1260')
@@ -450,7 +484,15 @@ describe('Config.vue network config projection', () => {
disable_ipv6: false,
hostname: 'host-edited',
proxy_cidrs: ['10.7.0.0/16', '172.17.0.0/16'],
vpn_portal_client_network_addr: '10.200.0.0',
vpn_portal_config: {
wireguard_listen: '[::]:23000',
wireguard_private_key: 'edited-private-key',
clients: [{
name: 'laptop-a',
virtual_ip: '10.1.2.20',
groups: ['ops', 'admin'],
}],
},
listener_urls: ['tcp://0.0.0.0:13010'],
dev_name: 'tun-edited',
mtu: 1260,
@@ -478,6 +520,15 @@ describe('Config.vue network config projection', () => {
listener_urls: ['tcp://0.0.0.0:13010'],
mtu: 1260,
instance_recv_bps_limit: '9007199254740993',
vpn_portal_config: {
wireguard_listen: '[::]:23000',
wireguard_private_key: 'edited-private-key',
clients: [{
name: 'laptop-a',
virtual_ip: '10.1.2.20',
groups: ['ops', 'admin'],
}],
},
port_forwards: [{
proto: 'tcp',
bind_ip: '127.0.0.1',
@@ -509,12 +560,13 @@ describe('Config.vue network config projection', () => {
}
const toggleButtons = wrapper.findAll('button[data-stub="toggle-button"]')
expect(toggleButtons).toHaveLength(CONFIG_TOGGLE_FIELDS.length)
expect(toggleButtons).toHaveLength(CONFIG_TOGGLE_FIELDS.length + 1)
for (const [index, field] of CONFIG_TOGGLE_FIELDS.entries()) {
const value = originalFlagValues.get(field)
expect(toggleButtons[index].attributes('aria-pressed'), `${field} should project into UI`)
const toggle = toggleButtons[index + 1]
expect(toggle.attributes('aria-pressed'), `${field} should project into UI`)
.toBe(String(value))
await toggleButtons[index].trigger('click')
await toggle.trigger('click')
await nextTick()
}
@@ -526,6 +578,55 @@ describe('Config.vue network config projection', () => {
}
})
it('uses VPN Portal config presence as the enable switch', async () => {
const config = DEFAULT_NETWORK_CONFIG()
const { curNetwork, wrapper } = mountConfig(config)
await nextTick()
const portalToggle = wrapper.findAll('button[data-stub="toggle-button"]')[0]
expect(portalToggle.attributes('aria-pressed')).toBe('false')
await portalToggle.trigger('click')
await nextTick()
expect(curNetwork.vpn_portal_config).toEqual({
wireguard_listen: '0.0.0.0:22022',
clients: [],
})
await portalToggle.trigger('click')
await nextTick()
expect(curNetwork.vpn_portal_config).toBeUndefined()
})
it('keeps each VPN Portal client row bound to the same client when reordered', async () => {
const config = makeConfig()
config.vpn_portal_config!.clients.push({
name: 'phone-b',
virtual_ip: '10.1.2.11',
groups: ['guests'],
})
const { curNetwork, wrapper } = mountConfig(config)
await nextTick()
const firstClient = curNetwork.vpn_portal_config!.clients[0]
const secondClient = curNetwork.vpn_portal_config!.clients[1]
const firstClientInput = input(wrapper, '#vpn_portal_client_name_0')
curNetwork.vpn_portal_config!.clients = [secondClient, firstClient]
await nextTick()
expect(input(wrapper, '#vpn_portal_client_name_1')).toBe(firstClientInput)
await setInput(wrapper, '#vpn_portal_client_name_1', 'phone-a-edited')
expect(firstClient.name).toBe('phone-a-edited')
expect(secondClient.name).toBe('phone-b')
})
it('keeps VPN Portal ACL group menus inside the management drawer', async () => {
const { wrapper } = mountConfig()
await nextTick()
expect(wrapper.find('#vpn_portal_client_groups_0').attributes('data-append-to')).toBe('self')
})
it('keeps uint64 input editable without losing large values', async () => {
const { curNetwork, wrapper } = mountConfig()
await nextTick()
@@ -9,7 +9,6 @@ import {
const BOOLEAN_CONFIG_FIELDS = [
'dhcp',
'enable_vpn_portal',
'advanced_settings',
'latency_first',
'use_smoltcp',
@@ -171,6 +170,7 @@ describe('RemoteManagement config save', () => {
generate_config: vi.fn(),
get_network_config: vi.fn(async () => cloneConfig(config)),
get_network_info: vi.fn(),
get_vpn_portal_info: vi.fn(),
get_network_metas: vi.fn(async (instanceIds: string[]) => ({
metas: Object.fromEntries(instanceIds.map((id) => [id, {
config_permission: 0xffffffff,
@@ -0,0 +1,160 @@
import { flushPromises, mount } from '@vue/test-utils'
import { describe, expect, it, vi } from 'vitest'
import { defineComponent, h } from 'vue'
import Status from '../src/components/Status.vue'
import { VpnPortalClientState, type NetworkInstance } from '../src/types/network'
vi.mock('vue-i18n', () => ({
useI18n: () => ({ t: (key: string) => key }),
}))
vi.mock('@vueuse/core', () => ({
useTimeAgo: () => '',
}))
vi.mock('../src/components/NetworkChart.vue', () => ({
default: defineComponent({ render: () => h('div') }),
}))
vi.mock('primevue', () => {
const PassThrough = defineComponent({
setup(_, { slots }) {
return () => h('div', slots.default?.())
},
})
const CardStub = defineComponent({
setup(_, { slots }) {
return () => h('div', [slots.title?.(), slots.content?.()])
},
})
const ButtonStub = defineComponent({
props: { label: String },
emits: ['click'],
setup(props, { emit }) {
return () => h('button', {
'data-label': props.label,
onClick: (event: MouseEvent) => emit('click', event),
}, props.label)
},
})
return {
Badge: PassThrough,
Button: ButtonStub,
Card: CardStub,
Chip: PassThrough,
Column: PassThrough,
DataTable: PassThrough,
Dialog: PassThrough,
Divider: PassThrough,
ScrollPanel: PassThrough,
Tag: PassThrough,
Timeline: PassThrough,
}
})
function runningInstance(): NetworkInstance {
return {
instance_id: '12345678-9abc-def0-fedc-ba9876543210',
running: true,
error_msg: '',
detail: {
dev_name: 'tun0',
running: true,
events: [],
routes: [],
peers: [],
peer_route_pairs: [],
my_node_info: {
virtual_ipv4: { address: { addr: 0x0a000001 }, network_length: 24 },
hostname: 'portal-node',
version: 'test',
ips: {
public_ipv4: { addr: 0 },
interface_ipv4s: [],
public_ipv6: { part1: 0, part2: 0, part3: 0, part4: 0 },
interface_ipv6s: [],
listeners: [],
},
stun_info: { udp_nat_type: 0, tcp_nat_type: 0, last_update_time: 0 },
listeners: [],
peer_id: 1,
},
},
}
}
describe('Status VPN Portal details', () => {
it('fetches client configs only when the user opens the dialog', async () => {
const getVpnPortalInfo = vi.fn(async () => ({
vpn_type: 'wireguard',
client_config: '',
connected_clients: [],
listener: '0.0.0.0:22022',
clients: [{
name: 'phone-a',
virtual_ip: '10.0.0.10',
groups: ['ops'],
state: VpnPortalClientState.ONLINE,
peer_id: 42,
endpoint: '203.0.113.5:51820',
tunnel_ip: '192.0.2.1',
client_config: '[Interface]\nPrivateKey = secret',
}],
}))
const wrapper = mount(Status, {
props: {
curNetworkInst: runningInstance(),
api: { get_vpn_portal_info: getVpnPortalInfo } as any,
},
global: {
directives: { tooltip: () => {} },
stubs: { HumanEvent: true },
},
})
try {
expect(getVpnPortalInfo).not.toHaveBeenCalled()
await wrapper.find('button[data-label="show_vpn_portal_config"]').trigger('click')
await flushPromises()
expect(getVpnPortalInfo).toHaveBeenCalledOnce()
expect(getVpnPortalInfo).toHaveBeenCalledWith('12345678-9abc-def0-fedc-ba9876543210')
expect(wrapper.text()).toContain('phone-a · 10.0.0.10')
expect(wrapper.text()).toContain('203.0.113.5:51820')
expect(wrapper.text()).toContain('PrivateKey = secret')
} finally {
wrapper.unmount()
}
})
it('renders the unconfigured portal sentinel as an empty state', async () => {
const getVpnPortalInfo = vi.fn(async () => ({
vpn_type: 'null',
client_config: '',
connected_clients: [],
clients: [],
}))
const wrapper = mount(Status, {
props: {
curNetworkInst: runningInstance(),
api: { get_vpn_portal_info: getVpnPortalInfo } as any,
},
global: {
directives: { tooltip: () => {} },
stubs: { HumanEvent: true },
},
})
try {
await wrapper.find('button[data-label="show_vpn_portal_config"]').trigger('click')
await flushPromises()
expect(wrapper.text()).toContain('vpn_portal_not_configured')
expect(wrapper.text()).not.toContain('vpn_portal_type: null')
} finally {
wrapper.unmount()
}
})
})
@@ -0,0 +1,21 @@
import { describe, expect, it } from 'vitest'
import { StrToUuid, UuidToStr } from '../src/modules/utils'
describe('UUID protobuf conversion', () => {
it('round-trips all four uint32 parts', () => {
const value = '12345678-9abc-def0-fedc-ba9876543210'
const protobuf = StrToUuid(value)
expect(protobuf).toEqual({
part1: 0x12345678,
part2: 0x9abcdef0,
part3: 0xfedcba98,
part4: 0x76543210,
})
expect(UuidToStr(protobuf)).toBe(value)
})
it('rejects malformed UUIDs', () => {
expect(() => StrToUuid('not-a-uuid')).toThrow('Invalid UUID')
})
})
+17
View File
@@ -220,6 +220,23 @@ class WebRemoteClient implements Api.RemoteClient {
const response = await this.client.get<any, Api.CollectNetworkInfoResponse>('/machines/' + this.machine_id + '/networks/info/' + inst_id);
return response.info?.map?.[inst_id];
}
async get_vpn_portal_info(inst_id: string): Promise<NetworkTypes.VpnPortalInfo | undefined> {
const response = await this.client.post<any, { vpn_portal_info?: NetworkTypes.VpnPortalInfo }>(
`/machines/${this.machine_id}/proxy-rpc`,
{
service_name: 'api.instance.VpnPortalRpcService',
method_name: 'get_vpn_portal_info',
payload: {
instance: {
id: Utils.StrToUuid(inst_id),
},
},
},
);
return response.vpn_portal_info
? NetworkTypes.normalizeVpnPortalInfo(response.vpn_portal_info)
: undefined;
}
async list_network_instance_ids(): Promise<Api.ListNetworkInstanceIdResponse> {
const response = await this.client.get<any, ListNetworkInstanceIdResponse>('/machines/' + this.machine_id + '/networks');
return response;
+3 -1
View File
@@ -179,6 +179,8 @@ network-interface = "2.0.5"
# for wireguard
boringtun = { package = "boringtun-easytier", version = "0.6.1", optional = true }
hkdf = { version = "0.12", optional = true }
sha2 = { version = "0.10", optional = true }
# for cli
tabled = "0.16"
@@ -348,7 +350,7 @@ full = [
"extended-services",
"tcp-hole-punch",
]
wireguard = ["vpn-portal", "dep:boringtun", "ring-crypto", "easytier-proto/wireguard"]
wireguard = ["vpn-portal", "dep:boringtun", "dep:hkdf", "dep:sha2", "ring-crypto", "easytier-proto/wireguard"]
quic = ["wrapped-transport", "easytier-core/proxy-packet", "dep:quinn", "dep:quinn-proto", "dep:seahash", "dep:rustls", "easytier-proto/quic"]
kcp = ["wrapped-transport", "easytier-core/proxy-packet", "dep:kcp-sys"]
mimalloc = ["dep:mimalloc"]
+11 -2
View File
@@ -100,8 +100,17 @@ core_clap:
en: "instance name to identify this vpn node in same machine"
zh-CN: "实例名称,用于在同一台机器上标识此VPN节点"
vpn_portal:
en: "url that defines the vpn portal, allow other vpn clients to connect. example: wg://0.0.0.0:11010/10.14.14.0/24, means the vpn portal is a wireguard server listening on vpn.example.com:11010, and the vpn client is in network of 10.14.14.0/24"
zh-CN: "定义VPN门户URL允许其他VPN客户端连接。示例:wg://0.0.0.0:11010/10.14.14.0/24,表示VPN门户是监听在vpn.example.com:11010的wireguard服务器,VPN客户端在10.14.14.0/24网络中"
en: "WireGuard VPN portal listener URL, for example wg://0.0.0.0:51820"
zh-CN: "WireGuard VPN 门户监听 URL例如 wg://0.0.0.0:51820"
vpn_portal_private_key:
en: "base64 WireGuard server private key (prefer ET_VPN_PORTAL_PRIVATE_KEY over command-line exposure)"
zh-CN: "Base64 WireGuard 服务端私钥(建议通过 ET_VPN_PORTAL_PRIVATE_KEY 传入,避免命令行暴露)"
vpn_portal_client:
en: "named VPN portal client in NAME=IP form; may be repeated"
zh-CN: "NAME=IP 格式的具名 VPN 门户客户端;可重复指定"
vpn_portal_client_group:
en: "VPN portal client group membership in NAME=GROUP form; may be repeated"
zh-CN: "NAME=GROUP 格式的 VPN 门户客户端组成员关系;可重复指定"
default_protocol:
en: "default protocol to use when connecting to peers"
zh-CN: "连接到对等节点时使用的默认协议"
-13
View File
@@ -88,7 +88,6 @@ pub struct GlobalCtx {
cached_ipv4: AtomicCell<Option<cidr::Ipv4Inet>>,
cached_ipv6: AtomicCell<Option<cidr::Ipv6Inet>>,
vpn_portal_cidr: AtomicCell<Option<cidr::Ipv4Cidr>>,
hostname: Mutex<String>,
tun_device_name: Mutex<Option<String>>,
@@ -188,13 +187,6 @@ impl GlobalCtx {
let ipv6 = runtime
.map(|runtime| Self::runtime_ipv6(&runtime.peer))
.unwrap_or_else(|| config_fs.get_ipv6());
let vpn_portal_cidr = runtime
.map(|runtime| runtime.peer.vpn_portal_cidr)
.unwrap_or_else(|| {
config_fs
.get_vpn_portal_config()
.map(|config| config.client_cidr)
});
if flags.enable_encryption && effective_encryption_uses_xor(&flags.encryption_algorithm) {
tracing::warn!("using insecure XOR because no AEAD encryption is configured");
}
@@ -210,7 +202,6 @@ impl GlobalCtx {
event_bus,
cached_ipv4: AtomicCell::new(ipv4),
cached_ipv6: AtomicCell::new(ipv6),
vpn_portal_cidr: AtomicCell::new(vpn_portal_cidr),
hostname: Mutex::new(hostname),
tun_device_name: Mutex::new(None),
@@ -326,10 +317,6 @@ impl GlobalCtx {
*self.hostname.lock().unwrap() = hostname;
}
pub fn get_vpn_portal_cidr(&self) -> Option<cidr::Ipv4Cidr> {
self.vpn_portal_cidr.load()
}
pub fn get_flags(&self) -> Flags {
self.flags.load().as_ref().clone()
}
+263 -19
View File
@@ -4,8 +4,8 @@ use crate::{
config::{
ConfigFileControl, ConfigLoader, ConsoleLoggerConfig, EncryptionAlgorithm,
FileLoggerConfig, LoggingConfigLoader, NetworkIdentity, PeerConfig, PortForwardConfig,
TomlConfigLoader, VpnPortalConfig, add_proxy_network_to_config, load_config_from_file,
load_toml_config_from_path, parse_mapped_listener_urls,
TomlConfigLoader, VpnPortalClientConfig, VpnPortalConfig, add_proxy_network_to_config,
load_config_from_file, load_toml_config_from_path, parse_mapped_listener_urls,
},
constants::EASYTIER_VERSION,
log,
@@ -281,6 +281,29 @@ struct NetworkOptions {
)]
vpn_portal: Option<String>,
#[arg(
long,
env = "ET_VPN_PORTAL_PRIVATE_KEY",
help = t!("core_clap.vpn_portal_private_key").to_string()
)]
vpn_portal_private_key: Option<String>,
#[arg(
long = "vpn-portal-client",
env = "ET_VPN_PORTAL_CLIENT",
value_delimiter = ',',
help = t!("core_clap.vpn_portal_client").to_string()
)]
vpn_portal_clients: Vec<String>,
#[arg(
long = "vpn-portal-client-group",
env = "ET_VPN_PORTAL_CLIENT_GROUP",
value_delimiter = ',',
help = t!("core_clap.vpn_portal_client_group").to_string()
)]
vpn_portal_client_groups: Vec<String>,
#[arg(
long,
env = "ET_DEFAULT_PROTOCOL",
@@ -864,6 +887,77 @@ impl Cli {
}
impl NetworkOptions {
fn parse_vpn_portal_listener(value: &str) -> anyhow::Result<SocketAddr> {
let url: url::Url = value
.parse()
.with_context(|| format!("failed to parse vpn portal url: {value}"))?;
if url.scheme() != "wg" {
anyhow::bail!("vpn portal URL must use the wg scheme: {value}");
}
if !url.path().is_empty() {
anyhow::bail!(
"legacy VPN portal CIDR paths are no longer supported; use wg://host:port and configure --vpn-portal-client NAME=IP"
);
}
if !url.username().is_empty()
|| url.password().is_some()
|| url.query().is_some()
|| url.fragment().is_some()
{
anyhow::bail!("vpn portal URL must have the form wg://host:port");
}
let host: IpAddr = url
.host_str()
.ok_or_else(|| anyhow::anyhow!("vpn portal url missing host"))?
.parse()
.with_context(|| "vpn portal listener host must be an IP address")?;
let port = url
.port()
.ok_or_else(|| anyhow::anyhow!("vpn portal url missing port"))?;
Ok(SocketAddr::new(host, port))
}
fn parse_vpn_portal_clients(&self) -> anyhow::Result<Vec<VpnPortalClientConfig>> {
let mut clients = self
.vpn_portal_clients
.iter()
.map(|value| {
let (name, virtual_ip) = value.split_once('=').ok_or_else(|| {
anyhow::anyhow!("invalid vpn portal client {value:?}; expected NAME=IP")
})?;
if name.is_empty() {
anyhow::bail!("vpn portal client name cannot be empty");
}
Ok(VpnPortalClientConfig {
name: name.to_owned(),
virtual_ip: virtual_ip.parse().with_context(|| {
format!("invalid virtual IP for vpn portal client {name}: {virtual_ip}")
})?,
groups: Vec::new(),
})
})
.collect::<anyhow::Result<Vec<_>>>()?;
for value in &self.vpn_portal_client_groups {
let (name, group) = value.split_once('=').ok_or_else(|| {
anyhow::anyhow!("invalid vpn portal client group {value:?}; expected NAME=GROUP")
})?;
if name.is_empty() || group.is_empty() {
anyhow::bail!("vpn portal client group name and group cannot be empty");
}
let client = clients
.iter_mut()
.find(|client| client.name == name)
.ok_or_else(|| {
anyhow::anyhow!("vpn portal client group references unknown CLI client: {name}")
})?;
client.groups.push(group.to_owned());
}
Ok(clients)
}
fn can_merge(
&self,
cfg: &TomlConfigLoader,
@@ -999,23 +1093,41 @@ impl NetworkOptions {
cfg.set_inst_name(inst_name.clone());
}
if let Some(vpn_portal) = self.vpn_portal.as_ref() {
let url: url::Url = vpn_portal
.parse()
.with_context(|| format!("failed to parse vpn portal url: {}", vpn_portal))?;
let host = url
.host_str()
.ok_or_else(|| anyhow::anyhow!("vpn portal url missing host"))?;
let port = url
.port()
.ok_or_else(|| anyhow::anyhow!("vpn portal url missing port"))?;
let client_cidr = url.path()[1..].parse().with_context(|| {
format!("failed to parse vpn portal client cidr: {}", url.path())
})?;
let wireguard_listen: SocketAddr = format!("{}:{}", host, port).parse().unwrap();
let has_vpn_portal_overrides = self.vpn_portal.is_some()
|| self.vpn_portal_private_key.is_some()
|| !self.vpn_portal_clients.is_empty()
|| !self.vpn_portal_client_groups.is_empty();
if has_vpn_portal_overrides {
if self.vpn_portal_clients.is_empty() && !self.vpn_portal_client_groups.is_empty() {
anyhow::bail!(
"--vpn-portal-client-group requires at least one --vpn-portal-client"
);
}
let existing = cfg.get_vpn_portal_config();
let wireguard_listen = match self.vpn_portal.as_deref() {
Some(value) => Self::parse_vpn_portal_listener(value)?,
None => existing
.as_ref()
.map(|portal| portal.wireguard_listen)
.ok_or_else(|| {
anyhow::anyhow!("--vpn-portal is required when no vpn_portal_config exists")
})?,
};
let wireguard_private_key = self
.vpn_portal_private_key
.clone()
.or_else(|| existing.as_ref()?.wireguard_private_key.clone());
let clients = if self.vpn_portal_clients.is_empty() {
existing.map_or_else(Vec::new, |portal| portal.clients)
} else {
self.parse_vpn_portal_clients()?
};
cfg.set_vpn_portal_config(VpnPortalConfig {
wireguard_listen,
client_cidr,
wireguard_private_key,
clients,
});
}
@@ -1497,7 +1609,7 @@ async fn run_main(cli: Cli) -> anyhow::Result<()> {
",
config_file,
control.permission,
cfg.dump()
cfg.dump_redacted()
);
manager.run_network_instance(cfg, control)?;
}
@@ -1514,7 +1626,7 @@ async fn run_main(cli: Cli) -> anyhow::Result<()> {
{}\n\
-----------------------------------\n\
",
cfg.dump()
cfg.dump_redacted()
);
manager.run_network_instance(cfg, ConfigFileControl::STATIC_CONFIG)?;
}
@@ -1826,4 +1938,136 @@ tcp_stun_servers = ["tcp.example.com:3478"]
assert_eq!(cfg.get_stun_servers_v6(), Some(Vec::new()));
assert_eq!(cfg.get_tcp_stun_servers(), Some(Vec::new()));
}
#[test]
fn vpn_portal_cli_uses_named_clients_and_preserves_unset_fields() {
let cfg = TomlConfigLoader::new_from_str(
r#"
[vpn_portal_config]
wireguard_listen = "127.0.0.1:51820"
wireguard_private_key = "existing-key"
[[vpn_portal_config.clients]]
name = "existing"
virtual_ip = "10.144.144.9"
"#,
)
.unwrap();
NetworkOptions {
vpn_portal: Some("wg://0.0.0.0:51821".to_owned()),
..Default::default()
}
.merge_into(&cfg)
.unwrap();
let preserved = cfg.get_vpn_portal_config().unwrap();
assert_eq!(preserved.wireguard_listen, "0.0.0.0:51821".parse().unwrap());
assert_eq!(
preserved.wireguard_private_key.as_deref(),
Some("existing-key")
);
assert_eq!(preserved.clients[0].name, "existing");
NetworkOptions {
vpn_portal_private_key: Some("replacement-key".to_owned()),
vpn_portal_clients: vec![
"alice=10.144.144.10".to_owned(),
"bob=10.144.144.11".to_owned(),
],
vpn_portal_client_groups: vec!["alice=staff".to_owned(), "alice=dev".to_owned()],
..Default::default()
}
.merge_into(&cfg)
.unwrap();
let replaced = cfg.get_vpn_portal_config().unwrap();
assert_eq!(replaced.wireguard_listen, "0.0.0.0:51821".parse().unwrap());
assert_eq!(
replaced.wireguard_private_key.as_deref(),
Some("replacement-key")
);
assert_eq!(
replaced
.clients
.iter()
.map(|client| client.name.as_str())
.collect::<Vec<_>>(),
vec!["alice", "bob"]
);
assert_eq!(
replaced.clients[0].groups,
vec!["staff".to_owned(), "dev".to_owned()]
);
assert!(replaced.clients[1].groups.is_empty());
}
#[test]
fn vpn_portal_cli_rejects_legacy_path_and_invalid_group_mapping() {
let error = NetworkOptions::parse_vpn_portal_listener("wg://0.0.0.0:51820/10.14.14.0/24")
.unwrap_err()
.to_string();
assert!(error.contains("legacy VPN portal CIDR"), "{error}");
let cfg = TomlConfigLoader::default();
let missing_clients = NetworkOptions {
vpn_portal: Some("wg://0.0.0.0:51820".to_owned()),
vpn_portal_client_groups: vec!["alice=staff".to_owned()],
..Default::default()
}
.merge_into(&cfg)
.unwrap_err()
.to_string();
assert!(
missing_clients.contains("requires at least one --vpn-portal-client"),
"{missing_clients}"
);
let unknown_client = NetworkOptions {
vpn_portal: Some("wg://0.0.0.0:51820".to_owned()),
vpn_portal_clients: vec!["alice=10.144.144.10".to_owned()],
vpn_portal_client_groups: vec!["bob=staff".to_owned()],
..Default::default()
}
.merge_into(&TomlConfigLoader::default())
.unwrap_err()
.to_string();
assert!(
unknown_client.contains("unknown CLI client: bob"),
"{unknown_client}"
);
}
#[test]
fn vpn_portal_cli_repeat_flags_use_singular_names() {
let cli = Cli::try_parse_from([
"easytier-core",
"--vpn-portal",
"wg://0.0.0.0:51820",
"--vpn-portal-private-key",
"private-key",
"--vpn-portal-client",
"alice=10.144.144.10",
"--vpn-portal-client",
"bob=10.144.144.11",
"--vpn-portal-client-group",
"alice=staff",
])
.unwrap();
assert_eq!(
cli.network_options.vpn_portal_clients,
vec![
"alice=10.144.144.10".to_owned(),
"bob=10.144.144.11".to_owned()
]
);
assert_eq!(
cli.network_options.vpn_portal_client_groups,
vec!["alice=staff".to_owned()]
);
assert_eq!(
cli.network_options.vpn_portal_private_key.as_deref(),
Some("private-key")
);
}
}
+27 -7
View File
@@ -2513,15 +2513,35 @@ impl<'a> CommandHandler<'a> {
self.print_results(&results, |resp| {
println!("portal_name: {}", resp.vpn_type);
if let Some(listener) = &resp.listener {
println!("listener: {listener}");
}
for client in &resp.clients {
let state = easytier_proto::api::instance::VpnPortalClientState::try_from(
client.state,
)
.map_or("UNKNOWN", |state| state.as_str_name());
println!(
r#"
############### client_config_start ###############
{}
############### client_config_end ###############
"#,
resp.client_config
"\nclient: {}\nvirtual_ip: {}\nstate: {}\npeer_id: {}\nendpoint: {}\ntunnel_ip: {}\ngroups: {}",
client.name,
client.virtual_ip,
state,
client
.peer_id
.map(|value| value.to_string())
.unwrap_or_else(|| "-".to_owned()),
client.endpoint.as_deref().unwrap_or("-"),
client.tunnel_ip.as_deref().unwrap_or("-"),
client.groups.join(", "),
);
println!("connected_clients:\n{:#?}", resp.connected_clients);
println!(
"############### client_config_start ###############\n{}############### client_config_end ###############",
client.client_config
);
if let Some(error) = &client.error {
println!("error: {error}");
}
}
Ok(())
})
}
+5 -6
View File
@@ -3,7 +3,7 @@ use std::sync::Arc;
#[cfg(feature = "wrapped-transport")]
use easytier_core::gateway::proxy::wrapped_transport::WrappedTransportEngines;
#[cfg(feature = "wireguard")]
use easytier_core::gateway::vpn_portal::VpnPortalHost;
use easytier_core::gateway::vpn_portal::PortalHost;
#[cfg(test)]
use easytier_core::host::packet::{HostPacket, PacketSink};
#[cfg(feature = "web-client")]
@@ -255,13 +255,12 @@ fn configure_runtime_core_host_adapters(
if host_config.vpn_portal_enabled {
use crate::common::config::ConfigLoader as _;
if let Some(config) = global_ctx.config.get_vpn_portal_config() {
adapters.vpn_portal = Some(crate::vpn_portal::wireguard::WireGuardPortalHost::new(
global_ctx.clone(),
global_ctx
.config
.get_vpn_portal_config()
.map(|config| config.wireguard_listen),
) as Arc<dyn VpnPortalHost>);
config,
) as Arc<dyn PortalHost>);
}
}
adapters
}
+359 -14
View File
@@ -7,6 +7,8 @@ use std::{
time::Duration,
};
#[cfg(feature = "wireguard")]
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD};
use easytier_core::{
connectivity::protocol::raw::TunnelDialer,
foundation::stats::{LabelSet, LabelType, MetricName, MetricSnapshot},
@@ -99,7 +101,13 @@ async fn set_foreign_network_refresh_interval(inst: &Instance, seconds: u64) {
}
#[cfg(feature = "wireguard")]
use crate::{common::config::VpnPortalConfig, vpn_portal::wireguard::get_wg_config_for_portal};
use crate::{
common::config::{VpnPortalClientConfig, VpnPortalConfig},
vpn_portal::wireguard::test_wireguard_keys,
};
#[cfg(feature = "wireguard")]
use easytier_core::gateway::vpn_portal::PortalClientState;
pub fn prepare_linux_namespaces() {
del_netns("net_a");
@@ -107,12 +115,14 @@ pub fn prepare_linux_namespaces() {
del_netns("net_c");
del_netns("net_d");
del_netns("net_e");
del_netns("net_f");
create_netns("net_a", "10.1.1.1/24", "fd11::1/64");
create_netns("net_b", "10.1.1.2/24", "fd11::2/64");
create_netns("net_c", "10.1.2.3/24", "fd12::3/64");
create_netns("net_d", "10.1.2.4/24", "fd12::4/64");
create_netns("net_e", "10.1.1.3/24", "fd11::3/64");
create_netns("net_f", "10.1.2.5/24", "fd12::5/64");
prepare_bridge("br_a");
prepare_bridge("br_b");
@@ -122,6 +132,7 @@ pub fn prepare_linux_namespaces() {
add_ns_to_bridge("br_a", "net_e");
add_ns_to_bridge("br_b", "net_c");
add_ns_to_bridge("br_b", "net_d");
add_ns_to_bridge("br_b", "net_f");
}
pub fn get_inst_config(
@@ -1676,7 +1687,17 @@ use defguard_wireguard_rs::{
InterfaceConfiguration, WGApi, WireguardInterfaceApi, host::Peer, key::Key, net::IpAddrMask,
};
fn wireguard_ifname(base: &str) -> String {
if cfg!(target_os = "linux") || cfg!(target_os = "freebsd") {
base.to_owned()
} else {
"utun3".into()
}
}
#[allow(clippy::too_many_arguments)]
fn run_wireguard_client(
ifname: &str,
endpoint: SocketAddr,
peer_public_key: Key,
client_private_key: Key,
@@ -1684,12 +1705,7 @@ fn run_wireguard_client(
client_ip: String,
) -> Result<(), Box<dyn std::error::Error>> {
// Create new API object for interface
let ifname: String = if cfg!(target_os = "linux") || cfg!(target_os = "freebsd") {
"wg0".into()
} else {
"utun3".into()
};
let wgapi = WGApi::new(ifname.clone(), false)?;
let wgapi = WGApi::new(ifname.to_owned(), false)?;
// create interface
wgapi.create_interface()?;
@@ -1707,7 +1723,7 @@ fn run_wireguard_client(
// interface configuration
let interface_config = InterfaceConfiguration {
name: ifname.clone(),
name: ifname.to_owned(),
prvkey: client_private_key.to_string(),
address: client_ip,
port: 12345,
@@ -1730,10 +1746,25 @@ pub async fn wireguard_vpn_portal(#[values(true, false)] test_v6: bool) {
let insts = init_three_node_ex(
"tcp",
|config| {
let identity = config.get_network_identity();
config.set_network_identity(NetworkIdentity::new(
identity.network_name,
"wireguard-portal-test".to_owned(),
));
if config.get_inst_name() == "inst1" {
config
.add_proxy_cidr("198.51.100.0/24".parse().unwrap(), None)
.unwrap();
}
if config.get_inst_name() == "inst3" {
config.set_vpn_portal_config(VpnPortalConfig {
wireguard_listen: "0.0.0.0:22121".parse().unwrap(),
client_cidr: "10.14.14.0/24".parse().unwrap(),
wireguard_private_key: Some(BASE64_STANDARD.encode([42u8; 32])),
clients: vec![VpnPortalClientConfig {
name: "test-client".to_owned(),
virtual_ip: "10.144.144.4".parse().unwrap(),
groups: Vec::new(),
}],
});
}
config
@@ -1756,13 +1787,31 @@ pub async fn wireguard_vpn_portal(#[values(true, false)] test_v6: bool) {
let net_ns = NetNS::new(Some("net_d".into()));
let _g = net_ns.guard();
let wg_cfg = get_wg_config_for_portal(&insts[2].get_global_ctx().get_network_identity());
let portal_config = insts[2]
.get_global_ctx()
.config
.get_vpn_portal_config()
.unwrap();
let portal_info = insts[2].get_core_instance().vpn_portal_info().await;
assert_eq!(portal_info.clients.len(), 1);
let client_info = portal_info
.clients
.iter()
.find(|client| client.name == "test-client")
.expect("configured client must be reported");
assert!(
client_info.client_config.contains("198.51.100.0/24"),
"client config must include remote proxy CIDRs"
);
let (server_public, client_private) =
test_wireguard_keys(&portal_config, "test-client").unwrap();
run_wireguard_client(
&wireguard_ifname("wg0"),
dst_socket_addr,
Key::try_from(wg_cfg.my_public_key()).unwrap(),
Key::try_from(wg_cfg.peer_secret_key()).unwrap(),
vec!["10.14.14.0/24".to_string(), "10.144.144.0/24".to_string()],
"10.14.14.2".to_string(),
Key::try_from(server_public.as_slice()).unwrap(),
Key::try_from(client_private.as_slice()).unwrap(),
vec!["10.144.144.0/24".to_string()],
"192.0.2.42".to_string(),
)
.unwrap();
@@ -1788,6 +1837,302 @@ pub async fn wireguard_vpn_portal(#[values(true, false)] test_v6: bool) {
drop_insts(insts).await;
}
#[cfg(feature = "wireguard")]
#[tokio::test]
#[serial_test::serial]
pub async fn wireguard_vpn_portal_multi_client() {
use rand::Rng as _;
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
let insts = init_three_node_ex(
"tcp",
|config| {
let identity = config.get_network_identity();
config.set_network_identity(NetworkIdentity::new(
identity.network_name,
"wireguard-portal-multi-client-test".to_owned(),
));
if config.get_inst_name() == "inst3" {
config.set_vpn_portal_config(VpnPortalConfig {
wireguard_listen: "0.0.0.0:22121".parse().unwrap(),
wireguard_private_key: Some(BASE64_STANDARD.encode([42u8; 32])),
clients: vec![
VpnPortalClientConfig {
name: "client-a".to_owned(),
virtual_ip: "10.144.144.4".parse().unwrap(),
groups: Vec::new(),
},
VpnPortalClientConfig {
name: "client-b".to_owned(),
virtual_ip: "10.144.144.5".parse().unwrap(),
groups: Vec::new(),
},
],
});
}
config
},
false,
)
.await;
let portal_config = insts[2]
.get_global_ctx()
.config
.get_vpn_portal_config()
.unwrap();
for (ns, client_name, tunnel_ip) in [
("net_d", "client-a", "192.0.2.42"),
("net_f", "client-b", "192.0.2.43"),
] {
let net_ns = NetNS::new(Some(ns.into()));
let _g = net_ns.guard();
let (server_public, client_private) =
test_wireguard_keys(&portal_config, client_name).unwrap();
run_wireguard_client(
&wireguard_ifname("wg0"),
"10.1.2.3:22121".parse().unwrap(),
Key::try_from(server_public.as_slice()).unwrap(),
Key::try_from(client_private.as_slice()).unwrap(),
vec!["10.144.144.0/24".to_string()],
tunnel_ip.to_string(),
)
.unwrap();
}
// 两个客户端各自 ping mesh 内节点
for ns in ["net_d", "net_f"] {
wait_for_condition(
|| async { ping_test(ns, "10.144.144.1", None).await },
Duration::from_secs(10),
)
.await;
wait_for_condition(
|| async { ping_test(ns, "10.144.144.2", None).await },
Duration::from_secs(10),
)
.await;
}
// 跨客户端互 ping 对方的虚拟 IP:一次流量同时覆盖源地址改写
// tunnel_ip -> virtual_ip)与目的地址改写(virtual_ip -> tunnel_ip),
// 回程再反向各执行一遍
wait_for_condition(
|| async { ping_test("net_d", "10.144.144.5", None).await },
Duration::from_secs(10),
)
.await;
wait_for_condition(
|| async { ping_test("net_f", "10.144.144.4", None).await },
Duration::from_secs(10),
)
.await;
// TCP 数据面:node1 侧看到的连接源地址必须是 client-a 的虚拟 IP
// 并做一段随机数据回环,覆盖 TCP 增量校验和改写路径
let mut buf = vec![0u8; 1024];
rand::thread_rng().fill(&mut buf[..]);
let expected = buf.clone();
let echo_task = tokio::spawn(async move {
let net_ns = NetNS::new(Some("net_a".into()));
let _g = net_ns.guard();
let socket = TcpListener::bind("0.0.0.0:22222").await.unwrap();
let (mut st, addr) = socket.accept().await.unwrap();
assert_eq!(addr.ip().to_string(), "10.144.144.4".to_string());
let mut rbuf = vec![0u8; 1024];
st.read_exact(&mut rbuf).await.unwrap();
assert_eq!(rbuf, expected);
st.write_all(&rbuf).await.unwrap();
});
{
let net_ns = NetNS::new(Some("net_d".into()));
let _g = net_ns.guard();
let mut stream = TcpStream::connect("10.144.144.1:22222").await.unwrap();
stream.write_all(&buf).await.unwrap();
let mut rbuf = vec![0u8; 1024];
stream.read_exact(&mut rbuf).await.unwrap();
assert_eq!(rbuf, buf);
}
echo_task.await.unwrap();
// portal 状态:两个客户端均在线,tunnel_ip 学习正确,peer_id 互不相同
let portal_info = insts[2].get_core_instance().vpn_portal_info().await;
assert_eq!(portal_info.clients.len(), 2);
let client_a = portal_info
.clients
.iter()
.find(|client| client.name == "client-a")
.expect("client-a must be reported");
let client_b = portal_info
.clients
.iter()
.find(|client| client.name == "client-b")
.expect("client-b must be reported");
for client in [client_a, client_b] {
assert_eq!(client.state, PortalClientState::Online);
assert!(client.peer_id.is_some());
}
assert_ne!(client_a.peer_id, client_b.peer_id);
assert_eq!(client_a.tunnel_ip, Some("192.0.2.42".parse().unwrap()));
assert_eq!(client_b.tunnel_ip, Some("192.0.2.43".parse().unwrap()));
drop_insts(insts).await;
}
#[cfg(feature = "wireguard")]
#[tokio::test]
#[serial_test::serial]
pub async fn wireguard_vpn_portal_client_roaming() {
let insts = init_three_node_ex(
"tcp",
|config| {
let identity = config.get_network_identity();
config.set_network_identity(NetworkIdentity::new(
identity.network_name,
"wireguard-portal-roaming-test".to_owned(),
));
if config.get_inst_name() == "inst3" {
config.set_vpn_portal_config(VpnPortalConfig {
wireguard_listen: "0.0.0.0:22121".parse().unwrap(),
wireguard_private_key: Some(BASE64_STANDARD.encode([42u8; 32])),
clients: vec![VpnPortalClientConfig {
name: "roaming-client".to_owned(),
virtual_ip: "10.144.144.4".parse().unwrap(),
groups: Vec::new(),
}],
});
}
config
},
false,
)
.await;
let portal_config = insts[2]
.get_global_ctx()
.config
.get_vpn_portal_config()
.unwrap();
{
let net_ns = NetNS::new(Some("net_d".into()));
let _g = net_ns.guard();
let (server_public, client_private) =
test_wireguard_keys(&portal_config, "roaming-client").unwrap();
run_wireguard_client(
&wireguard_ifname("wg0"),
"10.1.2.3:22121".parse().unwrap(),
Key::try_from(server_public.as_slice()).unwrap(),
Key::try_from(client_private.as_slice()).unwrap(),
vec!["10.144.144.0/24".to_string()],
"192.0.2.42".to_string(),
)
.unwrap();
}
// 客户端在 net_d(源地址 10.1.2.4)上线
wait_for_condition(
|| async { ping_test("net_d", "10.144.144.1", None).await },
Duration::from_secs(10),
)
.await;
let peer_id = {
let info = insts[2].get_core_instance().vpn_portal_info().await;
let client = info
.clients
.iter()
.find(|client| client.name == "roaming-client")
.expect("roaming client must be reported");
assert_eq!(client.state, PortalClientState::Online);
assert!(
client
.endpoint
.as_deref()
.is_some_and(|endpoint| endpoint.starts_with("10.1.2.4:")),
"unexpected endpoint before roaming: {:?}",
client.endpoint
);
client.peer_id
};
assert!(peer_id.is_some());
// 模拟客户端换网络:net_d 把地址从 10.1.2.4 换成 10.1.2.9。内核
// WireGuard 为 peer endpoint 缓存的源地址随旧地址一起失效,客户端
// 不重建 peer、不重新握手,直接用原 session 从新源继续发数据包,
// portal 应在数据路径上更新 endpoint
for args in [
vec![
"netns".to_owned(),
"exec".to_owned(),
"net_d".to_owned(),
"ip".to_owned(),
"addr".to_owned(),
"del".to_owned(),
"10.1.2.4/24".to_owned(),
"dev".to_owned(),
get_guest_veth_name("net_d").to_owned(),
],
vec![
"netns".to_owned(),
"exec".to_owned(),
"net_d".to_owned(),
"ip".to_owned(),
"addr".to_owned(),
"add".to_owned(),
"10.1.2.9/24".to_owned(),
"dev".to_owned(),
get_guest_veth_name("net_d").to_owned(),
],
] {
let ret = std::process::Command::new("ip")
.args(&args)
.output()
.unwrap();
assert!(
ret.status.success(),
"ip {args:?} failed: {}",
String::from_utf8_lossy(&ret.stderr)
);
}
// 驱动流量(ping 经内核 WireGuard 加密后从新源发出),portal 应在
// 同一个 peer 上更新 endpointpeer_id 不变说明是同代漫游,客户端没有掉线重连
wait_for_condition(
|| async {
ping_test("net_d", "10.144.144.1", None).await;
let info = insts[2].get_core_instance().vpn_portal_info().await;
info.clients.iter().any(|client| {
client.name == "roaming-client"
&& client.state == PortalClientState::Online
&& client.peer_id == peer_id
&& client
.endpoint
.as_deref()
.is_some_and(|endpoint| endpoint.starts_with("10.1.2.9:"))
})
},
Duration::from_secs(20),
)
.await;
// 漫游后连通性保持
wait_for_condition(
|| async { ping_test("net_d", "10.144.144.1", None).await },
Duration::from_secs(10),
)
.await;
wait_for_condition(
|| async { ping_test("net_d", "10.144.144.3", None).await },
Duration::from_secs(10),
)
.await;
drop_insts(insts).await;
}
#[cfg(feature = "wireguard")]
#[rstest::rstest]
#[tokio::test]
+367 -90
View File
@@ -1,129 +1,406 @@
//! Native WireGuard Adapter for the VPN portal.
//!
//! One UDP socket and one shared MAC/cookie limiter demultiplex all configured
//! public keys. Per-client slots keep rekeys and endpoint roaming within one
//! authenticated generation; Core receives a new session only after the prior
//! generation has expired and been detached atomically.
mod engine;
use std::{
fmt,
net::{Ipv6Addr, SocketAddr, SocketAddrV6},
sync::Arc,
};
use anyhow::Context;
use base64::{Engine, prelude::BASE64_STANDARD};
use anyhow::Context as _;
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD};
use boringtun::x25519::{PublicKey, StaticSecret};
use easytier_core::{
gateway::vpn_portal::{VpnPortalClientConfigPlan, VpnPortalHost, VpnPortalListener},
socket::SocketListener,
config::toml::VpnPortalConfig,
gateway::vpn_portal::{PortalClientConfigPlan, PortalHost, PortalListener, PortalSession},
socket::{
ListenerConnectionCounter, NetNamespace, SocketContext, SocketListener,
udp::{UdpBindOptions, VirtualUdpSocket, VirtualUdpSocketFactory},
},
};
use hkdf::Hkdf;
use sha2::Sha256;
use tokio::{sync::mpsc, task::JoinSet};
use crate::{
common::{config::NetworkIdentity, global_ctx::ArcGlobalCtx},
tunnel::wireguard::{WgConfig, WgTunnelListener},
common::global_ctx::ArcGlobalCtx,
host_runtime::{NativeHostRuntime, native_host_runtime},
socket::udp::RuntimeUdpSocket,
};
pub(crate) fn get_wg_config_for_portal(nid: &NetworkIdentity) -> WgConfig {
let key_seed = format!(
"{}{}",
nid.network_name,
nid.network_secret.as_ref().unwrap_or(&String::new())
);
WgConfig::new_for_portal(&key_seed, &key_seed)
use self::engine::{DerivedClient, PortalEngine};
struct WireGuardPortalListener {
url: url::Url,
sockets: Vec<Arc<RuntimeUdpSocket>>,
receiver: mpsc::UnboundedReceiver<PortalSession>,
engine: Arc<PortalEngine>,
tasks: JoinSet<anyhow::Result<()>>,
listened: bool,
}
fn listener_endpoint(listener_url: &url::Url) -> &str {
&listener_url[url::Position::BeforeHost..url::Position::AfterPort]
}
pub struct WireGuardPortalHost {
global_ctx: ArcGlobalCtx,
wg_config: WgConfig,
listener_addr: Option<SocketAddr>,
}
impl WireGuardPortalHost {
pub fn new(global_ctx: ArcGlobalCtx, listener_addr: Option<SocketAddr>) -> Arc<Self> {
Arc::new(Self {
wg_config: get_wg_config_for_portal(&global_ctx.get_network_identity()),
global_ctx,
listener_addr,
})
}
async fn start_listener(&self, listener_addr: SocketAddr) -> anyhow::Result<VpnPortalListener> {
let mut listener_url = url::Url::parse("wg://0.0.0.0:0").unwrap();
listener_url.set_port(Some(listener_addr.port())).unwrap();
listener_url.set_ip_host(listener_addr.ip()).unwrap();
let mut listener = WgTunnelListener::new(listener_url, self.wg_config.clone());
{
let _guard = self.global_ctx.net_ns.guard();
listener
.listen()
.await
.context("failed to start WireGuard VPN portal listener")?;
}
Ok(Box::new(listener))
impl fmt::Debug for WireGuardPortalListener {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("WireGuardPortalListener")
.field("url", &self.url)
.finish_non_exhaustive()
}
}
#[async_trait::async_trait]
impl VpnPortalHost for WireGuardPortalHost {
async fn start_listeners(&self) -> anyhow::Result<Vec<VpnPortalListener>> {
let listener_addr = self.listener_addr.context("VPN portal config is not set")?;
let mut listeners = vec![self.start_listener(listener_addr).await?];
if let SocketAddr::V4(v4) = listener_addr
&& v4.ip().is_unspecified()
&& let Ok(listener) = self
.start_listener(SocketAddr::V6(SocketAddrV6::new(
Ipv6Addr::UNSPECIFIED,
v4.port(),
0,
0,
)))
.await
{
listeners.push(listener);
impl SocketListener for WireGuardPortalListener {
type Accepted = PortalSession;
async fn listen(&mut self) -> anyhow::Result<()> {
if self.listened {
return Ok(());
}
Ok(listeners)
self.listened = true;
for socket in &self.sockets {
let socket = socket.clone();
let engine = self.engine.clone();
self.tasks.spawn(async move {
loop {
let datagram = socket
.recv_session_datagram()
.await
.context("WireGuard portal UDP receive failed")?;
engine
.handle_datagram(socket.clone(), datagram.remote_addr, &datagram.payload)
.await;
}
});
}
let engine = self.engine.clone();
self.tasks.spawn(async move {
engine.run_timers().await;
Ok(())
});
Ok(())
}
async fn accept(&mut self) -> anyhow::Result<Self::Accepted> {
let tasks_active = !self.tasks.is_empty();
tokio::select! {
biased;
task = self.tasks.join_next(), if tasks_active => {
let error = match task {
Some(Ok(Err(error))) => {
anyhow::anyhow!("WireGuard portal reader stopped: {error:#}")
}
Some(Ok(Ok(()))) => {
anyhow::anyhow!("WireGuard portal task stopped unexpectedly")
}
Some(Err(error)) => {
anyhow::anyhow!("WireGuard portal task failed: {error}")
}
None => anyhow::anyhow!("WireGuard portal listener stopped"),
};
self.engine.cancel();
self.tasks.abort_all();
Err(error)
}
session = self.receiver.recv() => {
session.ok_or_else(|| anyhow::anyhow!("WireGuard portal listener stopped"))
}
}
}
fn local_url(&self) -> url::Url {
self.url.clone()
}
fn connection_counter(&self) -> Arc<dyn ListenerConnectionCounter> {
Arc::new(PortalConnectionCounter(self.engine.clone()))
}
}
impl Drop for WireGuardPortalListener {
fn drop(&mut self) {
self.engine.cancel();
self.tasks.abort_all();
}
}
struct PortalConnectionCounter(Arc<PortalEngine>);
impl fmt::Debug for PortalConnectionCounter {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("PortalConnectionCounter").finish()
}
}
impl ListenerConnectionCounter for PortalConnectionCounter {
fn get(&self) -> Option<u32> {
Some(self.0.connection_count())
}
}
pub struct WireGuardPortalHost {
global_ctx: ArcGlobalCtx,
config: VpnPortalConfig,
setup: Result<WireGuardPortalSetup, String>,
}
struct WireGuardPortalSetup {
server_private: [u8; 32],
server_public: PublicKey,
clients: Vec<DerivedClient>,
}
impl WireGuardPortalHost {
pub fn new(global_ctx: ArcGlobalCtx, config: VpnPortalConfig) -> Arc<Self> {
let setup = (|| -> anyhow::Result<_> {
let (master, server_private) = portal_master_and_server_key(&config)?;
let server_public = PublicKey::from(&StaticSecret::from(server_private));
let mut clients = Vec::with_capacity(config.clients.len());
for client in &config.clients {
let wireguard_private =
derive_named_key(&master, b"wireguard-client", &client.name)?;
let identity_private_key =
derive_named_key(&master, b"attached-noise", &client.name)?;
clients.push(DerivedClient {
config: client.clone(),
wireguard_private,
wireguard_public: PublicKey::from(&StaticSecret::from(wireguard_private)),
identity_private_key,
});
}
Ok(WireGuardPortalSetup {
server_private,
server_public,
clients,
})
})()
.map_err(|error| error.to_string());
Arc::new(Self {
global_ctx,
config,
setup,
})
}
async fn bind_socket(
&self,
runtime: &NativeHostRuntime,
address: SocketAddr,
only_v6: bool,
) -> anyhow::Result<Arc<RuntimeUdpSocket>> {
let context = SocketContext::default()
.with_socket_mark(self.global_ctx.get_flags().socket_mark)
.with_netns(self.global_ctx.net_ns.name().map(NetNamespace::new));
runtime
.bind_udp(
UdpBindOptions::port_bound_listener(address)
.with_context(context)
.with_only_v6(only_v6),
)
.await
}
}
fn secondary_ipv6_bind_address(address: SocketAddr, primary_port: u16) -> Option<SocketAddr> {
let SocketAddr::V4(address) = address else {
return None;
};
address
.ip()
.is_unspecified()
.then(|| SocketAddr::V6(SocketAddrV6::new(Ipv6Addr::UNSPECIFIED, primary_port, 0, 0)))
}
#[async_trait::async_trait]
impl PortalHost for WireGuardPortalHost {
async fn start_listeners(&self) -> anyhow::Result<Vec<PortalListener>> {
let setup = self
.setup
.as_ref()
.map_err(|error| anyhow::anyhow!(error.clone()))?;
let runtime = native_host_runtime();
let primary = self
.bind_socket(&runtime, self.config.wireguard_listen, false)
.await
.context("failed to bind WireGuard VPN portal")?;
let local = primary.local_addr()?;
let mut sockets = vec![primary.clone()];
if let Some(v6_address) =
secondary_ipv6_bind_address(self.config.wireguard_listen, local.port())
&& let Ok(v6) = self.bind_socket(&runtime, v6_address, true).await
{
sockets.push(v6);
}
let url = url::Url::parse(&format!("wg://{local}"))?;
let (accepted, receiver) = mpsc::unbounded_channel();
let engine = PortalEngine::new(setup.server_private, setup.clients.clone(), accepted);
Ok(vec![Box::new(WireGuardPortalListener {
url,
sockets,
receiver,
engine,
tasks: JoinSet::new(),
listened: false,
})])
}
fn name(&self) -> String {
"wireguard".to_owned()
}
fn render_client_config(&self, plan: &VpnPortalClientConfigPlan) -> String {
let listener_addr = listener_endpoint(&plan.listener_url);
fn render_client_config(&self, plan: &PortalClientConfigPlan) -> String {
let client = self
.setup
.as_ref()
.expect("client config is rendered only after successful startup")
.clients
.iter()
.find(|client| client.config.name == plan.name)
.expect("Core only renders configured clients");
let endpoint = &plan.listener_url[url::Position::BeforeHost..url::Position::AfterPort];
format!(
r#"
[Interface]
PrivateKey = {peer_secret_key}
Address = {address} # should assign an ip from this cidr manually
[Peer]
PublicKey = {my_public_key}
AllowedIPs = {allowed_ips}
Endpoint = {listener_addr} # should be the public ip(or domain) of the vpn server
PersistentKeepalive = 25
"#,
peer_secret_key = BASE64_STANDARD.encode(self.wg_config.peer_secret_key()),
my_public_key = BASE64_STANDARD.encode(self.wg_config.my_public_key()),
listener_addr = listener_addr,
allowed_ips = plan.allowed_ips.join(","),
address = plan.client_cidr.first_address().to_string() + "/32",
"[Interface]\nPrivateKey = {}\nAddress = {}/32\n\n[Peer]\nPublicKey = {}\nAllowedIPs = {}\nEndpoint = {} # replace wildcard with the public address\nPersistentKeepalive = 25\n",
BASE64_STANDARD.encode(client.wireguard_private),
plan.address,
BASE64_STANDARD.encode(
self.setup
.as_ref()
.expect("client config is rendered only after successful startup")
.server_public
.as_bytes()
),
plan.allowed_ips.join(", "),
endpoint,
)
}
fn not_started_client_config(&self) -> String {
"ERROR: Wireguard VPN Portal Not Started".to_owned()
}
fn portal_master_and_server_key(config: &VpnPortalConfig) -> anyhow::Result<([u8; 32], [u8; 32])> {
let encoded = config
.wireguard_private_key
.as_deref()
.filter(|key| !key.is_empty())
.ok_or_else(|| anyhow::anyhow!("WireGuard portal requires a dedicated private key"))?;
let decoded = BASE64_STANDARD
.decode(encoded)
.context("invalid base64 WireGuard portal private key")?;
let key: [u8; 32] = decoded
.try_into()
.map_err(|_| anyhow::anyhow!("WireGuard portal private key must be 32 bytes"))?;
Ok((key, key))
}
fn derive_named_key(master: &[u8; 32], domain: &[u8], name: &str) -> anyhow::Result<[u8; 32]> {
let mut context = Vec::new();
append_label(&mut context, domain);
append_label(&mut context, name.as_bytes());
let hkdf = Hkdf::<Sha256>::new(Some(b"easytier/wireguard-portal/v1"), master);
let mut output = [0u8; 32];
hkdf.expand(&context, &mut output)
.map_err(|_| anyhow::anyhow!("failed to derive portal client key"))?;
Ok(output)
}
fn append_label(output: &mut Vec<u8>, label: &[u8]) {
output.extend_from_slice(&(label.len() as u32).to_be_bytes());
output.extend_from_slice(label);
}
#[cfg(test)]
pub(crate) fn test_wireguard_keys(
config: &VpnPortalConfig,
client_name: &str,
) -> anyhow::Result<([u8; 32], [u8; 32])> {
let (master, server_private) = portal_master_and_server_key(config)?;
let client_private = derive_named_key(&master, b"wireguard-client", client_name)?;
let server_public = *PublicKey::from(&StaticSecret::from(server_private)).as_bytes();
Ok((server_public, client_private))
}
#[cfg(test)]
mod tests {
use super::listener_endpoint;
use super::*;
#[test]
fn listener_endpoint_uses_the_active_listener_url() {
fn named_keys_are_domain_separated_and_stable() {
let master = [7; 32];
let client = derive_named_key(&master, b"wireguard-client", "laptop").unwrap();
assert_eq!(
listener_endpoint(&"wg://192.0.2.10:51820".parse().unwrap()),
"192.0.2.10:51820"
client,
derive_named_key(&master, b"wireguard-client", "laptop").unwrap()
);
assert_eq!(
listener_endpoint(&"wg://[2001:db8::10]:51820".parse().unwrap()),
"[2001:db8::10]:51820"
assert_ne!(
client,
derive_named_key(&master, b"attached-noise", "laptop").unwrap()
);
assert_ne!(
client,
derive_named_key(&master, b"wireguard-client", "phone").unwrap()
);
}
#[test]
fn explicit_server_key_is_the_derivation_master() {
let key = [9; 32];
let config = VpnPortalConfig {
wireguard_listen: "127.0.0.1:51820".parse().unwrap(),
wireguard_private_key: Some(BASE64_STANDARD.encode(key)),
clients: Vec::new(),
};
assert_eq!(portal_master_and_server_key(&config).unwrap(), (key, key));
}
#[test]
fn portal_key_has_no_network_secret_fallback() {
let config = VpnPortalConfig {
wireguard_listen: "127.0.0.1:51820".parse().unwrap(),
wireguard_private_key: None,
clients: Vec::new(),
};
assert!(
portal_master_and_server_key(&config)
.unwrap_err()
.to_string()
.contains("dedicated private key")
);
}
#[test]
fn ipv6_wildcard_reuses_primary_ephemeral_port() {
assert_eq!(
secondary_ipv6_bind_address("0.0.0.0:0".parse().unwrap(), 43123),
Some("[::]:43123".parse().unwrap())
);
assert_eq!(
secondary_ipv6_bind_address("127.0.0.1:0".parse().unwrap(), 43123),
None
);
}
#[tokio::test]
async fn reader_failure_terminates_listener_accept() {
let (accepted, receiver) = mpsc::unbounded_channel();
let engine = PortalEngine::new([11; 32], Vec::new(), accepted);
let mut tasks: JoinSet<anyhow::Result<()>> = JoinSet::new();
tasks.spawn(async { anyhow::bail!("UDP receive failed") });
let mut listener = WireGuardPortalListener {
url: "wg://127.0.0.1:51820".parse().unwrap(),
sockets: Vec::new(),
receiver,
engine,
tasks,
listened: true,
};
let error = tokio::time::timeout(std::time::Duration::from_secs(1), listener.accept())
.await
.expect("listener accept ignored reader failure")
.unwrap_err();
assert!(error.to_string().contains("UDP receive failed"));
}
}
+435
View File
@@ -0,0 +1,435 @@
//! Shared WireGuard packet engine for named portal clients.
use atomic_shim::AtomicU64;
use std::{
collections::HashMap,
net::SocketAddr,
sync::{Arc, atomic::Ordering},
time::Duration,
};
use boringtun::{
noise::{
Packet, Tunn, TunnResult, errors::WireGuardError, handshake::parse_handshake_anon,
rate_limiter::RateLimiter,
},
x25519::{PublicKey, StaticSecret},
};
use easytier_core::{
config::toml::VpnPortalClientConfig, gateway::vpn_portal::PortalSession,
socket::udp::VirtualUdpSocket,
};
use tokio::{
sync::{Mutex, mpsc, watch},
task::JoinSet,
};
use tokio_util::sync::CancellationToken;
use crate::socket::udp::RuntimeUdpSocket;
const MIN_WIREGUARD_PACKET_CAPACITY: usize = 148;
// We pre-verify through this shared limiter and BoringTun verifies again inside
// each Tunn. Doubling the threshold preserves the intended 100 datagrams/s
// transition to cookies while retaining the upstream security ordering.
const DOUBLE_VERIFY_HANDSHAKE_LIMIT: u64 = 200;
const TIMER_INTERVAL: Duration = Duration::from_millis(250);
const PORTAL_PACKET_CAPACITY: usize = 128;
#[derive(Clone)]
pub(super) struct DerivedClient {
pub(super) config: VpnPortalClientConfig,
pub(super) wireguard_private: [u8; 32],
pub(super) wireguard_public: PublicKey,
pub(super) identity_private_key: [u8; 32],
}
struct PortalChannels {
endpoint: watch::Receiver<String>,
from_client: mpsc::Receiver<Vec<u8>>,
to_client: mpsc::Sender<Vec<u8>>,
}
struct ClientSession {
generation: u64,
endpoint: Option<Endpoint>,
endpoint_updates: watch::Sender<String>,
tunnel: Tunn,
from_client: mpsc::Sender<Vec<u8>>,
portal_channels: Option<PortalChannels>,
drain_capacity: usize,
tasks: JoinSet<()>,
}
struct ClientSlot {
client: DerivedClient,
index: u32,
next_generation: AtomicU64,
session: Mutex<Option<ClientSession>>,
}
#[derive(Clone)]
struct Endpoint {
socket: Arc<RuntimeUdpSocket>,
remote: SocketAddr,
}
impl ClientSession {
fn update_endpoint(&mut self, socket: Arc<RuntimeUdpSocket>, remote: SocketAddr) {
let changed = self
.endpoint
.as_ref()
.is_none_or(|endpoint| endpoint.remote != remote);
self.endpoint = Some(Endpoint { socket, remote });
if changed {
self.endpoint_updates.send_replace(remote.to_string());
}
}
}
pub(super) struct PortalEngine {
server_private: StaticSecret,
server_public: PublicKey,
rate_limiter: Arc<RateLimiter>,
by_public_key: HashMap<[u8; 32], Arc<ClientSlot>>,
by_index: HashMap<u32, Arc<ClientSlot>>,
accepted: mpsc::UnboundedSender<PortalSession>,
cancel: CancellationToken,
}
impl PortalEngine {
pub(super) fn new(
server_private: [u8; 32],
clients: Vec<DerivedClient>,
accepted: mpsc::UnboundedSender<PortalSession>,
) -> Arc<Self> {
let server_private = StaticSecret::from(server_private);
let server_public = PublicKey::from(&server_private);
let mut by_public_key = HashMap::with_capacity(clients.len());
let mut by_index = HashMap::with_capacity(clients.len());
for (offset, client) in clients.into_iter().enumerate() {
let index = u32::try_from(offset + 1).expect("client limit is below u32");
let public = *client.wireguard_public.as_bytes();
let slot = Arc::new(ClientSlot {
client,
index,
next_generation: AtomicU64::new(1),
session: Mutex::new(None),
});
by_public_key.insert(public, slot.clone());
by_index.insert(index, slot);
}
Arc::new(Self {
server_private,
server_public,
rate_limiter: Arc::new(RateLimiter::new(
&server_public,
DOUBLE_VERIFY_HANDSHAKE_LIMIT,
)),
by_public_key,
by_index,
accepted,
cancel: CancellationToken::new(),
})
}
pub(super) fn cancel(&self) {
self.cancel.cancel();
}
pub(super) fn connection_count(&self) -> u32 {
self.by_index
.values()
.filter(|slot| {
slot.session.try_lock().is_ok_and(|guard| {
guard
.as_ref()
.is_some_and(|session| session.portal_channels.is_none())
})
})
.count() as u32
}
pub(super) async fn handle_datagram(
self: &Arc<Self>,
socket: Arc<RuntimeUdpSocket>,
remote: SocketAddr,
datagram: &[u8],
) {
let mut cookie = [0u8; 148];
let parsed = match self
.rate_limiter
.verify_packet(Some(remote.ip()), datagram, &mut cookie)
{
Ok(packet) => packet,
Err(TunnResult::WriteToNetwork(reply)) => {
let _ = socket.send_to(reply, remote).await;
return;
}
Err(_) => return,
};
let slot = match &parsed {
Packet::HandshakeInit(init) => {
parse_handshake_anon(&self.server_private, &self.server_public, init)
.ok()
.and_then(|handshake| {
self.by_public_key
.get(&handshake.peer_static_public)
.cloned()
})
}
Packet::HandshakeResponse(response) => self.slot_by_receiver(response.receiver_idx),
Packet::PacketCookieReply(reply) => self.slot_by_receiver(reply.receiver_idx),
Packet::PacketData(data) => self.slot_by_receiver(data.receiver_idx),
};
let Some(slot) = slot else { return };
let mut session = slot.session.lock().await;
if session.is_none() {
if !matches!(parsed, Packet::HandshakeInit(_)) {
return;
}
*session = Some(self.new_session(&slot, socket.clone(), remote));
}
let current = session.as_mut().expect("created above");
let is_data = matches!(&parsed, Packet::PacketData(_));
let is_handshake_response = matches!(&parsed, Packet::HandshakeResponse(_));
// The shared pre-verification establishes the correct upstream order.
// Tunn::decapsulate performs a second MAC/cookie check because the
// dependency's verified-dispatch method is not public. Size the first
// output to the datagram: unauthenticated transport packets must not
// amplify a tiny allocation into a full-size IP buffer.
let mut output = vec![0u8; datagram.len().max(MIN_WIREGUARD_PACKET_CAPACITY)];
let mut result = current
.tunnel
.decapsulate(Some(remote.ip()), datagram, &mut output);
let mut first_result = true;
loop {
match result {
TunnResult::Done => {
if is_data {
current.update_endpoint(socket.clone(), remote);
self.activate_client(&slot, current);
}
current.drain_capacity = MIN_WIREGUARD_PACKET_CAPACITY;
break;
}
TunnResult::Err(WireGuardError::ConnectionExpired) => {
let expired = session.take();
drop(session);
Self::retire_session(expired);
return;
}
TunnResult::Err(_) => break,
TunnResult::WriteToNetwork(packet) => {
if (first_result && is_handshake_response && is_transport_data_packet(packet))
|| is_handshake_response_packet(packet)
{
current.update_endpoint(socket.clone(), remote);
}
let _ = socket.send_to(packet, remote).await;
// BoringTun queues a Core packet while it establishes a
// session. Its contract requires empty decapsulate calls
// after every network write until Done releases that queue.
first_result = false;
if output.len() < current.drain_capacity {
output.resize(current.drain_capacity, 0);
}
result = current.tunnel.decapsulate(None, &[], &mut output);
}
TunnResult::WriteToTunnelV4(packet, _) => {
current.update_endpoint(socket.clone(), remote);
self.activate_client(&slot, current);
match current.from_client.try_send(packet.to_vec()) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => {
tracing::debug!(
client = %slot.client.config.name,
"dropping WireGuard packet because the client queue is full"
);
}
Err(mpsc::error::TrySendError::Closed(_)) => {
let generation = current.generation;
drop(session);
self.expire_if_current(slot, generation).await;
return;
}
}
break;
}
TunnResult::WriteToTunnelV6(_, _) => {
// Portal traffic is deliberately IPv4-only.
break;
}
}
}
}
fn activate_client(&self, slot: &ClientSlot, session: &mut ClientSession) {
let Some(channels) = session.portal_channels.take() else {
return;
};
let _ = self.accepted.send(PortalSession {
client_name: slot.client.config.name.clone(),
endpoint: channels.endpoint,
identity_private_key: slot.client.identity_private_key,
from_client: channels.from_client,
to_client: channels.to_client,
});
}
fn slot_by_receiver(&self, receiver: u32) -> Option<Arc<ClientSlot>> {
self.by_index.get(&(receiver >> 8)).cloned()
}
fn new_session(
self: &Arc<Self>,
slot: &Arc<ClientSlot>,
socket: Arc<RuntimeUdpSocket>,
remote: SocketAddr,
) -> ClientSession {
let generation = slot.next_generation.fetch_add(1, Ordering::Relaxed);
let (from_client, portal_from_client) = mpsc::channel(PORTAL_PACKET_CAPACITY);
let (portal_to_client, mut to_client) = mpsc::channel::<Vec<u8>>(PORTAL_PACKET_CAPACITY);
let (endpoint_updates, portal_endpoint) = watch::channel(remote.to_string());
let engine = Arc::downgrade(self);
let slot_for_task = Arc::downgrade(slot);
let mut tasks = JoinSet::new();
tasks.spawn(async move {
while let Some(payload) = to_client.recv().await {
let Some(engine) = engine.upgrade() else {
return;
};
let Some(slot) = slot_for_task.upgrade() else {
return;
};
engine
.encapsulate_for_client(&slot, generation, &payload)
.await;
}
if let (Some(engine), Some(slot)) = (engine.upgrade(), slot_for_task.upgrade()) {
engine.expire_if_current(slot, generation).await;
}
});
ClientSession {
generation,
endpoint: Some(Endpoint { socket, remote }),
endpoint_updates,
tunnel: Tunn::new(
self.server_private.clone(),
slot.client.wireguard_public,
None,
None,
slot.index,
Some(self.rate_limiter.clone()),
),
from_client,
portal_channels: Some(PortalChannels {
endpoint: portal_endpoint,
from_client: portal_from_client,
to_client: portal_to_client,
}),
drain_capacity: MIN_WIREGUARD_PACKET_CAPACITY,
tasks,
}
}
async fn encapsulate_for_client(
self: &Arc<Self>,
slot: &Arc<ClientSlot>,
generation: u64,
payload: &[u8],
) {
let mut output = vec![0u8; payload.len().saturating_add(148).max(148)];
let mut guard = slot.session.lock().await;
let Some(session) = guard
.as_mut()
.filter(|session| session.generation == generation)
else {
return;
};
match session.tunnel.encapsulate(payload, &mut output) {
TunnResult::WriteToNetwork(packet) => {
if is_handshake_initiation(packet) {
session.drain_capacity =
session.drain_capacity.max(payload.len().saturating_add(32));
}
if let Some(endpoint) = session.endpoint.clone() {
let _ = endpoint.socket.send_to(packet, endpoint.remote).await;
}
}
TunnResult::Done => {
session.drain_capacity =
session.drain_capacity.max(payload.len().saturating_add(32));
}
TunnResult::Err(WireGuardError::ConnectionExpired) => {
drop(guard);
self.expire_if_current(slot.clone(), generation).await;
}
_ => {}
}
}
async fn expire_if_current(self: &Arc<Self>, slot: Arc<ClientSlot>, generation: u64) {
let expired = {
let mut guard = slot.session.lock().await;
if guard
.as_ref()
.is_some_and(|session| session.generation == generation)
{
guard.take()
} else {
None
}
};
Self::retire_session(expired);
}
fn retire_session(expired: Option<ClientSession>) {
if let Some(mut expired) = expired {
expired.tasks.abort_all();
// Dropping the ring sink atomically disconnects the matching Core
// generation. A newer generation, if any, owns a different ring.
}
}
pub(super) async fn run_timers(self: Arc<Self>) {
let mut interval = tokio::time::interval(TIMER_INTERVAL);
loop {
tokio::select! {
_ = self.cancel.cancelled() => return,
_ = interval.tick() => {}
}
self.rate_limiter.reset_count();
for slot in self.by_index.values() {
let mut output = [0u8; 148];
let mut guard = slot.session.lock().await;
let Some(session) = guard.as_mut() else {
continue;
};
match session.tunnel.update_timers(&mut output) {
TunnResult::WriteToNetwork(packet) => {
if let Some(endpoint) = session.endpoint.clone() {
let _ = endpoint.socket.send_to(packet, endpoint.remote).await;
}
}
TunnResult::Err(WireGuardError::ConnectionExpired) => {
let expired = guard.take();
drop(guard);
Self::retire_session(expired);
}
_ => {}
}
}
}
}
}
fn is_handshake_initiation(packet: &[u8]) -> bool {
packet.len() == 148 && packet.get(..4) == Some(&1u32.to_le_bytes())
}
fn is_handshake_response_packet(packet: &[u8]) -> bool {
packet.len() == 92 && packet.get(..4) == Some(&2u32.to_le_bytes())
}
fn is_transport_data_packet(packet: &[u8]) -> bool {
packet.len() >= 32 && packet.get(..4) == Some(&4u32.to_le_bytes())
}