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
+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()
}
+116 -16
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()
.await
.ok_or(anyhow::anyhow!("recv_packet_from_chan failed")),
Err(TryRecvError::Disconnected) => Err(anyhow::anyhow!("recv_packet_from_chan failed")),
}
recv_packet_envelope_from_chan(packet_recv_chan_receiver)
.await
.map(|envelope| envelope.packet)
}
#[async_trait::async_trait]
+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());