mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-04 01:55:41 +00:00
refactor(core): separate portable core from native runtime (#2451)
Create easytier-core as the portable owner of configuration, connectivity, tunnels, peer and routing state, gateways, management, the data plane, and instance lifecycle. Keep operating-system integration, native protocol engines, process startup, and presentation in easytier behind explicit Host capability adapters. Create easytier-proto to own schemas, generated RPC types, descriptors, and feature-scoped protocol slices. Remove runtime protobuf reflection from core while preserving unknown route-peer fields across forwarding. Normalize instance construction through CoreInstance, CoreHostAdapters, CoreProcessRuntime, and InstanceManager. Make the runtime config store the only authoritative mutable configuration after startup. Move the portable TCP/UDP data plane into core and extract a generic OperationBroker for completion, cancellation, disposal, and capacity accounting. Expose the session-based FFI v2 completion API and keep the WASI guest ABI, wire schemas, and adapters with core. Migrate CLI, GUI, web, FFI, Android JNI, OHOS, uptime, and mobile consumers to the shared manager and core state. Add explicit user/web config ownership and revision-aware web reconciliation. Preserve configuration, wire, and management behavior while fixing regressions discovered by the full platform and integration matrix: - inherit advertised relay capabilities in foreign networks; - refresh OSPF peer state immediately after runtime config changes; - restore CLI GlobalCtx event output without forcing GUI logging; - retain legacy encryption names and standalone RPC tunnel metadata; - restore ICMP host composition and fragmented UDP handling; - use portable 64-bit atomics on 32-bit MIPS targets; and - retain discarded operations until late cancellation completes. Validate the refactor across 45 GitHub checks, including Linux, macOS, Windows, FreeBSD, web, GUI, Android, OHOS, feature profiles, and three-node and subnet-proxy integration tests. BREAKING CHANGE: internal Rust module paths are not preserved. Legacy native data-plane APIs are replaced by the session-based FFI v2 API. The dedicated Android data-plane wrapper is removed.
This commit is contained in:
@@ -0,0 +1,650 @@
|
||||
use std::net::{Ipv4Addr, Ipv6Addr};
|
||||
use std::sync::atomic::Ordering;
|
||||
use std::{
|
||||
net::IpAddr,
|
||||
sync::{Arc, Mutex, atomic::AtomicBool},
|
||||
};
|
||||
|
||||
use arc_swap::ArcSwap;
|
||||
use dashmap::DashMap;
|
||||
use easytier_proto::acl::{Acl, AclStats, Action, ChainType, Protocol};
|
||||
use quanta::Instant;
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
|
||||
use crate::{
|
||||
packet::{PacketType, ZCPacket},
|
||||
peers::acl::processor::{AclProcessor, AclResult, AclStatKey, AclStatType, PacketInfo},
|
||||
};
|
||||
|
||||
const IP_PROTO_ICMP: u8 = 1;
|
||||
const IP_PROTO_TCP: u8 = 6;
|
||||
const IP_PROTO_UDP: u8 = 17;
|
||||
const IP_PROTO_ICMPV6: u8 = 58;
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct ParsedIpPacket<'a> {
|
||||
src_ip: IpAddr,
|
||||
dst_ip: IpAddr,
|
||||
protocol: u8,
|
||||
transport_payload: &'a [u8],
|
||||
}
|
||||
|
||||
fn parse_ip_packet(payload: &[u8]) -> Option<ParsedIpPacket<'_>> {
|
||||
let version = payload.first()? >> 4;
|
||||
match version {
|
||||
4 => parse_ipv4_packet(payload),
|
||||
6 => parse_ipv6_packet(payload),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_ipv4_packet(payload: &[u8]) -> Option<ParsedIpPacket<'_>> {
|
||||
if payload.len() < 20 {
|
||||
return None;
|
||||
}
|
||||
let header_len = usize::from(payload[0] & 0x0f) * 4;
|
||||
let options_len = header_len.saturating_sub(20);
|
||||
let payload_offset = 20 + options_len;
|
||||
let payload_start = payload_offset.min(payload.len());
|
||||
let total_length = usize::from(u16::from_be_bytes([payload[2], payload[3]]));
|
||||
let payload_len = total_length.saturating_sub(header_len);
|
||||
let payload_end = payload_start.saturating_add(payload_len).min(payload.len());
|
||||
|
||||
Some(ParsedIpPacket {
|
||||
src_ip: IpAddr::V4(Ipv4Addr::new(
|
||||
payload[12],
|
||||
payload[13],
|
||||
payload[14],
|
||||
payload[15],
|
||||
)),
|
||||
dst_ip: IpAddr::V4(Ipv4Addr::new(
|
||||
payload[16],
|
||||
payload[17],
|
||||
payload[18],
|
||||
payload[19],
|
||||
)),
|
||||
protocol: payload[9],
|
||||
transport_payload: &payload[payload_start..payload_end],
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_ipv6_packet(payload: &[u8]) -> Option<ParsedIpPacket<'_>> {
|
||||
if payload.len() < 40 {
|
||||
return None;
|
||||
}
|
||||
let payload_len = usize::from(u16::from_be_bytes([payload[4], payload[5]]));
|
||||
let payload_end = 40usize.saturating_add(payload_len).min(payload.len());
|
||||
|
||||
Some(ParsedIpPacket {
|
||||
src_ip: IpAddr::V6(Ipv6Addr::from(<[u8; 16]>::try_from(&payload[8..24]).ok()?)),
|
||||
dst_ip: IpAddr::V6(Ipv6Addr::from(<[u8; 16]>::try_from(&payload[24..40]).ok()?)),
|
||||
protocol: payload[6],
|
||||
transport_payload: &payload[40..payload_end],
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_transport_ports(protocol: u8, payload: &[u8]) -> Option<(Option<u16>, Option<u16>)> {
|
||||
let min_len = match protocol {
|
||||
IP_PROTO_TCP => 20,
|
||||
IP_PROTO_UDP => 8,
|
||||
_ => return Some((None, None)),
|
||||
};
|
||||
if payload.len() < min_len {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some((
|
||||
Some(u16::from_be_bytes([payload[0], payload[1]])),
|
||||
Some(u16::from_be_bytes([payload[2], payload[3]])),
|
||||
))
|
||||
}
|
||||
|
||||
fn acl_protocol(protocol: u8) -> Protocol {
|
||||
match protocol {
|
||||
IP_PROTO_TCP => Protocol::Tcp,
|
||||
IP_PROTO_UDP => Protocol::Udp,
|
||||
IP_PROTO_ICMP => Protocol::Icmp,
|
||||
IP_PROTO_ICMPV6 => Protocol::IcmPv6,
|
||||
_ => Protocol::Unspecified,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Eq, PartialEq, Hash)]
|
||||
struct OutboundAllowRecord {
|
||||
src_ip: IpAddr,
|
||||
dst_ip: IpAddr,
|
||||
src_port: Option<u16>,
|
||||
dst_port: Option<u16>,
|
||||
protocol: Protocol,
|
||||
}
|
||||
|
||||
impl OutboundAllowRecord {
|
||||
fn new_from_inbound_packet(p: &PacketInfo) -> Self {
|
||||
Self {
|
||||
src_ip: p.src_ip,
|
||||
dst_ip: p.dst_ip,
|
||||
src_port: p.src_port,
|
||||
dst_port: p.dst_port,
|
||||
protocol: p.protocol,
|
||||
}
|
||||
}
|
||||
|
||||
fn new_from_outbound_packet(p: &PacketInfo) -> Self {
|
||||
Self {
|
||||
src_ip: p.dst_ip,
|
||||
dst_ip: p.src_ip,
|
||||
src_port: p.dst_port,
|
||||
dst_port: p.src_port,
|
||||
protocol: p.protocol,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// ACL filter that can be inserted into the packet processing pipeline
|
||||
/// Optimized with lock-free hot reloading via atomic processor replacement
|
||||
pub struct AclFilter {
|
||||
// Use ArcSwap for lock-free atomic replacement during hot reload
|
||||
acl_processor: ArcSwap<AclProcessor>,
|
||||
acl_enabled: Arc<AtomicBool>,
|
||||
|
||||
// Track allowed outbound packets and automatically allow their corresponding inbound response
|
||||
// packets, even if they would normally be dropped by ACL rules
|
||||
outbound_allow_records: Arc<DashMap<OutboundAllowRecord, Instant>>,
|
||||
#[allow(dead_code)]
|
||||
clean_task: Mutex<Option<AbortOnDropHandle<()>>>,
|
||||
}
|
||||
|
||||
impl Default for AclFilter {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl AclFilter {
|
||||
pub fn new() -> Self {
|
||||
let outbound_allow_records = Arc::new(DashMap::new());
|
||||
let record_clone = outbound_allow_records.clone();
|
||||
Self {
|
||||
acl_processor: ArcSwap::from(Arc::new(AclProcessor::new(Acl::default()))),
|
||||
acl_enabled: Arc::new(AtomicBool::new(false)),
|
||||
outbound_allow_records,
|
||||
clean_task: Mutex::new(Some(AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
let max_life = std::time::Duration::from_secs(30);
|
||||
loop {
|
||||
record_clone.retain(|_, v| v.elapsed() < max_life);
|
||||
crate::foundation::time::sleep(std::time::Duration::from_secs(30)).await;
|
||||
}
|
||||
})))),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn stop_cleanup_task(&self) {
|
||||
let task = self.clean_task.lock().unwrap().take();
|
||||
if let Some(task) = task {
|
||||
task.abort();
|
||||
let _ = task.await;
|
||||
}
|
||||
}
|
||||
|
||||
/// Hot reload ACL rules by creating a new processor instance
|
||||
/// Preserves connection tracking and rate limiting state across reloads
|
||||
/// Now lock-free and doesn't require &mut self!
|
||||
pub fn reload_rules(&self, acl_config: Option<&Acl>) {
|
||||
self.outbound_allow_records.clear();
|
||||
|
||||
let Some(acl_config) = acl_config else {
|
||||
self.acl_enabled.store(false, Ordering::Relaxed);
|
||||
return;
|
||||
};
|
||||
|
||||
// Get current processor to extract shared state
|
||||
let current_processor = self.acl_processor.load();
|
||||
let (conn_track, rate_limiters, stats) = current_processor.get_shared_state();
|
||||
|
||||
// Create new processor with preserved state
|
||||
let new_processor = AclProcessor::new_with_shared_state(
|
||||
acl_config.clone(),
|
||||
Some(conn_track),
|
||||
Some(rate_limiters),
|
||||
Some(stats),
|
||||
);
|
||||
|
||||
// Atomic replacement - this is completely lock-free!
|
||||
self.acl_processor.store(Arc::new(new_processor));
|
||||
self.acl_enabled.store(true, Ordering::Relaxed);
|
||||
|
||||
tracing::info!("ACL rules hot reloaded with preserved state (lock-free)");
|
||||
}
|
||||
|
||||
/// Get current processor for processing packets
|
||||
pub fn get_processor(&self) -> Arc<AclProcessor> {
|
||||
self.acl_processor.load_full()
|
||||
}
|
||||
|
||||
pub fn get_stats(&self) -> AclStats {
|
||||
let processor = self.get_processor();
|
||||
let global_stats = processor.get_stats();
|
||||
let (conn_track, _, _) = processor.get_shared_state();
|
||||
let rules_stats = processor.get_rules_stats();
|
||||
|
||||
AclStats {
|
||||
global: global_stats.into_iter().collect(),
|
||||
conn_track: conn_track.iter().map(|x| *x.value()).collect(),
|
||||
rules: rules_stats,
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract packet information for ACL processing
|
||||
fn extract_packet_info(
|
||||
&self,
|
||||
packet: &ZCPacket,
|
||||
route: &(dyn crate::peers::route::Route + Send + Sync + 'static),
|
||||
) -> Option<PacketInfo> {
|
||||
let payload = packet.payload();
|
||||
|
||||
let parsed = parse_ip_packet(payload)?;
|
||||
let (src_port, dst_port) =
|
||||
parse_transport_ports(parsed.protocol, parsed.transport_payload)?;
|
||||
let acl_protocol = acl_protocol(parsed.protocol);
|
||||
|
||||
let src_groups = packet
|
||||
.get_src_peer_id()
|
||||
.map(|peer_id| route.get_peer_groups(peer_id))
|
||||
.unwrap_or_else(|| Arc::new(Vec::new()));
|
||||
let dst_groups = packet
|
||||
.get_dst_peer_id()
|
||||
.map(|peer_id| route.get_peer_groups(peer_id))
|
||||
.unwrap_or_else(|| Arc::new(Vec::new()));
|
||||
|
||||
Some(PacketInfo {
|
||||
src_ip: parsed.src_ip,
|
||||
dst_ip: parsed.dst_ip,
|
||||
src_port,
|
||||
dst_port,
|
||||
protocol: acl_protocol,
|
||||
packet_size: payload.len(),
|
||||
src_groups,
|
||||
dst_groups,
|
||||
})
|
||||
}
|
||||
|
||||
/// Process ACL result and log if needed
|
||||
pub fn handle_acl_result(
|
||||
&self,
|
||||
result: &AclResult,
|
||||
packet_info: &PacketInfo,
|
||||
chain_type: ChainType,
|
||||
processor: &AclProcessor,
|
||||
) {
|
||||
if result.should_log
|
||||
&& let Some(ref log_context) = result.log_context
|
||||
{
|
||||
let log_message = log_context.to_message();
|
||||
tracing::info!(
|
||||
src_ip = %packet_info.src_ip,
|
||||
dst_ip = %packet_info.dst_ip,
|
||||
src_port = packet_info.src_port,
|
||||
dst_port = packet_info.dst_port,
|
||||
src_group = packet_info.src_groups.join(","),
|
||||
dst_group = packet_info.dst_groups.join(","),
|
||||
protocol = ?packet_info.protocol,
|
||||
action = ?result.action,
|
||||
rule = result.matched_rule_str().as_deref().unwrap_or("unknown"),
|
||||
chain_type = ?chain_type,
|
||||
"ACL: {}", log_message
|
||||
);
|
||||
}
|
||||
|
||||
// Update global statistics in the ACL processor
|
||||
match result.action {
|
||||
Action::Allow => {
|
||||
processor.increment_stat(AclStatKey::PacketsAllowed);
|
||||
processor.increment_stat(AclStatKey::from_chain_and_action(
|
||||
chain_type,
|
||||
AclStatType::Allowed,
|
||||
));
|
||||
tracing::trace!("ACL: Packet allowed");
|
||||
}
|
||||
Action::Drop => {
|
||||
processor.increment_stat(AclStatKey::PacketsDropped);
|
||||
processor.increment_stat(AclStatKey::from_chain_and_action(
|
||||
chain_type,
|
||||
AclStatType::Dropped,
|
||||
));
|
||||
tracing::debug!("ACL: Packet dropped");
|
||||
}
|
||||
Action::Noop => {
|
||||
processor.increment_stat(AclStatKey::PacketsNoop);
|
||||
processor.increment_stat(AclStatKey::from_chain_and_action(
|
||||
chain_type,
|
||||
AclStatType::Noop,
|
||||
));
|
||||
tracing::trace!("ACL: No operation");
|
||||
}
|
||||
}
|
||||
|
||||
// Track total packets processed per chain
|
||||
processor.increment_stat(AclStatKey::from_chain_and_action(
|
||||
chain_type,
|
||||
AclStatType::Total,
|
||||
));
|
||||
processor.increment_stat(AclStatKey::PacketsTotal);
|
||||
}
|
||||
|
||||
fn classify_chain_type(
|
||||
is_in: bool,
|
||||
packet_info: &PacketInfo,
|
||||
my_ipv4: Option<Ipv4Addr>,
|
||||
is_local_ipv6: impl Fn(Ipv6Addr) -> bool,
|
||||
) -> ChainType {
|
||||
if !is_in {
|
||||
return ChainType::Outbound;
|
||||
}
|
||||
|
||||
let is_local_dst = packet_info.dst_ip == my_ipv4.unwrap_or(Ipv4Addr::UNSPECIFIED)
|
||||
|| matches!(packet_info.dst_ip, IpAddr::V6(dst) if is_local_ipv6(dst));
|
||||
|
||||
if is_local_dst {
|
||||
ChainType::Inbound
|
||||
} else {
|
||||
ChainType::Forward
|
||||
}
|
||||
}
|
||||
|
||||
/// Common ACL processing logic
|
||||
pub fn process_packet_with_acl(
|
||||
&self,
|
||||
packet: &ZCPacket,
|
||||
is_in: bool,
|
||||
my_ipv4: Option<Ipv4Addr>,
|
||||
is_local_ipv6: impl Fn(Ipv6Addr) -> bool,
|
||||
route: &(dyn crate::peers::route::Route + Send + Sync + 'static),
|
||||
) -> bool {
|
||||
if !self.acl_enabled.load(Ordering::Relaxed) {
|
||||
return true;
|
||||
}
|
||||
|
||||
if packet.peer_manager_header().unwrap().packet_type != PacketType::Data as u8 {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Extract packet information
|
||||
let packet_info = match self.extract_packet_info(packet, route) {
|
||||
Some(info) => info,
|
||||
None => {
|
||||
tracing::warn!(
|
||||
"Failed to extract packet info from {:?} packet, header: {:?}",
|
||||
if is_in { "inbound" } else { "outbound" },
|
||||
packet.peer_manager_header()
|
||||
);
|
||||
// allow all unknown packets
|
||||
return true;
|
||||
}
|
||||
};
|
||||
|
||||
let chain_type = Self::classify_chain_type(is_in, &packet_info, my_ipv4, is_local_ipv6);
|
||||
|
||||
// Get current processor atomically
|
||||
let processor = self.get_processor();
|
||||
|
||||
// Process through ACL rules
|
||||
let acl_result = processor.process_packet(&packet_info, chain_type);
|
||||
|
||||
self.handle_acl_result(&acl_result, &packet_info, chain_type, &processor);
|
||||
|
||||
// Check if packet should be allowed
|
||||
match acl_result.action {
|
||||
Action::Allow | Action::Noop => {
|
||||
if matches!(chain_type, ChainType::Outbound) {
|
||||
self.outbound_allow_records.insert(
|
||||
OutboundAllowRecord::new_from_outbound_packet(&packet_info),
|
||||
Instant::now(),
|
||||
);
|
||||
}
|
||||
true
|
||||
}
|
||||
Action::Drop => {
|
||||
if is_in {
|
||||
let record = OutboundAllowRecord::new_from_inbound_packet(&packet_info);
|
||||
let entry = self.outbound_allow_records.entry(record);
|
||||
if let dashmap::Entry::Occupied(mut entry) = entry {
|
||||
entry.insert(Instant::now());
|
||||
tracing::trace!(
|
||||
"ACL: Allowing {:?} packet from {} to {} because of existing allow record, chain_type: {:?}",
|
||||
packet_info.protocol,
|
||||
packet_info.src_ip,
|
||||
packet_info.dst_ip,
|
||||
chain_type,
|
||||
);
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
tracing::trace!(
|
||||
"ACL: Dropping {:?} packet from {} to {}, chain_type: {:?}",
|
||||
packet_info.protocol,
|
||||
packet_info.src_ip,
|
||||
packet_info.dst_ip,
|
||||
chain_type,
|
||||
);
|
||||
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::{
|
||||
net::{IpAddr, Ipv4Addr, Ipv6Addr},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use quanta::Instant;
|
||||
|
||||
use easytier_proto::acl::{Acl, ChainType, Protocol};
|
||||
|
||||
use crate::peers::acl::processor::PacketInfo;
|
||||
|
||||
use super::{
|
||||
AclFilter, IP_PROTO_ICMP, IP_PROTO_TCP, IP_PROTO_UDP, OutboundAllowRecord, acl_protocol,
|
||||
parse_ip_packet, parse_transport_ports,
|
||||
};
|
||||
|
||||
impl AclFilter {
|
||||
pub(crate) fn cleanup_task_is_stopped(&self) -> bool {
|
||||
self.clean_task.lock().unwrap().is_none()
|
||||
}
|
||||
}
|
||||
|
||||
fn packet_info(dst_ip: IpAddr) -> PacketInfo {
|
||||
PacketInfo {
|
||||
src_ip: IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)),
|
||||
dst_ip,
|
||||
src_port: Some(1234),
|
||||
dst_port: Some(80),
|
||||
protocol: Protocol::Tcp,
|
||||
packet_size: 64,
|
||||
src_groups: Arc::new(Vec::new()),
|
||||
dst_groups: Arc::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ipv4_tcp_packet_extracts_addrs_and_ports() {
|
||||
let mut packet = vec![0u8; 40];
|
||||
packet[0] = 0x45;
|
||||
packet[2..4].copy_from_slice(&40u16.to_be_bytes());
|
||||
packet[9] = IP_PROTO_TCP;
|
||||
packet[12..16].copy_from_slice(&[10, 0, 0, 1]);
|
||||
packet[16..20].copy_from_slice(&[10, 0, 0, 2]);
|
||||
packet[20..22].copy_from_slice(&1234u16.to_be_bytes());
|
||||
packet[22..24].copy_from_slice(&80u16.to_be_bytes());
|
||||
|
||||
let parsed = parse_ip_packet(&packet).unwrap();
|
||||
let (src_port, dst_port) =
|
||||
parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap();
|
||||
|
||||
assert_eq!(parsed.src_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)));
|
||||
assert_eq!(parsed.dst_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2)));
|
||||
assert_eq!(acl_protocol(parsed.protocol), Protocol::Tcp);
|
||||
assert_eq!(src_port, Some(1234));
|
||||
assert_eq!(dst_port, Some(80));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ipv6_udp_packet_extracts_addrs_and_ports() {
|
||||
let src: Ipv6Addr = "2001:db8::1".parse().unwrap();
|
||||
let dst: Ipv6Addr = "2001:db8::2".parse().unwrap();
|
||||
let mut packet = vec![0u8; 48];
|
||||
packet[0] = 0x60;
|
||||
packet[4..6].copy_from_slice(&8u16.to_be_bytes());
|
||||
packet[6] = IP_PROTO_UDP;
|
||||
packet[8..24].copy_from_slice(&src.octets());
|
||||
packet[24..40].copy_from_slice(&dst.octets());
|
||||
packet[40..42].copy_from_slice(&5353u16.to_be_bytes());
|
||||
packet[42..44].copy_from_slice(&53u16.to_be_bytes());
|
||||
|
||||
let parsed = parse_ip_packet(&packet).unwrap();
|
||||
let (src_port, dst_port) =
|
||||
parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap();
|
||||
|
||||
assert_eq!(parsed.src_ip, IpAddr::V6(src));
|
||||
assert_eq!(parsed.dst_ip, IpAddr::V6(dst));
|
||||
assert_eq!(acl_protocol(parsed.protocol), Protocol::Udp);
|
||||
assert_eq!(src_port, Some(5353));
|
||||
assert_eq!(dst_port, Some(53));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ipv4_uses_declared_total_length() {
|
||||
let mut packet = vec![0u8; 40];
|
||||
packet[0] = 0x45;
|
||||
packet[2..4].copy_from_slice(&24u16.to_be_bytes());
|
||||
packet[9] = IP_PROTO_TCP;
|
||||
packet[12..16].copy_from_slice(&[10, 0, 0, 1]);
|
||||
packet[16..20].copy_from_slice(&[10, 0, 0, 2]);
|
||||
packet[20..22].copy_from_slice(&1234u16.to_be_bytes());
|
||||
packet[22..24].copy_from_slice(&80u16.to_be_bytes());
|
||||
|
||||
let parsed = parse_ip_packet(&packet).unwrap();
|
||||
|
||||
assert_eq!(parsed.transport_payload.len(), 4);
|
||||
assert!(parse_transport_ports(parsed.protocol, parsed.transport_payload).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ipv6_uses_declared_payload_length() {
|
||||
let mut packet = vec![0u8; 48];
|
||||
packet[0] = 0x60;
|
||||
packet[4..6].copy_from_slice(&4u16.to_be_bytes());
|
||||
packet[6] = IP_PROTO_UDP;
|
||||
packet[40..42].copy_from_slice(&5353u16.to_be_bytes());
|
||||
packet[42..44].copy_from_slice(&53u16.to_be_bytes());
|
||||
|
||||
let parsed = parse_ip_packet(&packet).unwrap();
|
||||
|
||||
assert_eq!(parsed.transport_payload.len(), 4);
|
||||
assert!(parse_transport_ports(parsed.protocol, parsed.transport_payload).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ipv4_keeps_pnet_ihl_less_than_five_behavior() {
|
||||
let mut packet = vec![0u8; 40];
|
||||
packet[0] = 0x44;
|
||||
packet[2..4].copy_from_slice(&40u16.to_be_bytes());
|
||||
packet[9] = IP_PROTO_TCP;
|
||||
packet[20..22].copy_from_slice(&1234u16.to_be_bytes());
|
||||
packet[22..24].copy_from_slice(&80u16.to_be_bytes());
|
||||
|
||||
let parsed = parse_ip_packet(&packet).unwrap();
|
||||
let (src_port, dst_port) =
|
||||
parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap();
|
||||
|
||||
assert_eq!(parsed.transport_payload.len(), 20);
|
||||
assert_eq!(src_port, Some(1234));
|
||||
assert_eq!(dst_port, Some(80));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_ipv4_keeps_pnet_truncated_options_behavior() {
|
||||
let mut packet = vec![0u8; 20];
|
||||
packet[0] = 0x4f;
|
||||
packet[2..4].copy_from_slice(&60u16.to_be_bytes());
|
||||
packet[9] = IP_PROTO_ICMP;
|
||||
packet[12..16].copy_from_slice(&[10, 0, 0, 1]);
|
||||
packet[16..20].copy_from_slice(&[10, 0, 0, 2]);
|
||||
|
||||
let parsed = parse_ip_packet(&packet).unwrap();
|
||||
let (src_port, dst_port) =
|
||||
parse_transport_ports(parsed.protocol, parsed.transport_payload).unwrap();
|
||||
|
||||
assert_eq!(parsed.src_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)));
|
||||
assert_eq!(parsed.dst_ip, IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2)));
|
||||
assert_eq!(acl_protocol(parsed.protocol), Protocol::Icmp);
|
||||
assert!(parsed.transport_payload.is_empty());
|
||||
assert_eq!(src_port, None);
|
||||
assert_eq!(dst_port, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classify_chain_type_treats_public_ipv6_lease_as_inbound() {
|
||||
let leased_ipv6 = Ipv6Addr::new(0x2001, 0xdb8, 0x100, 0, 0, 0, 0, 0x123);
|
||||
let packet_info = packet_info(IpAddr::V6(leased_ipv6));
|
||||
|
||||
let chain =
|
||||
AclFilter::classify_chain_type(true, &packet_info, None, |ip| ip == leased_ipv6);
|
||||
|
||||
assert_eq!(chain, ChainType::Inbound);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classify_chain_type_keeps_non_local_ipv6_as_forward() {
|
||||
let leased_ipv6 = Ipv6Addr::new(0x2001, 0xdb8, 0x100, 0, 0, 0, 0, 0x123);
|
||||
let packet_info = packet_info(IpAddr::V6(Ipv6Addr::new(
|
||||
0x2001, 0xdb8, 0xffff, 2, 0, 0, 0, 0x100,
|
||||
)));
|
||||
|
||||
let chain =
|
||||
AclFilter::classify_chain_type(true, &packet_info, None, |ip| ip == leased_ipv6);
|
||||
|
||||
assert_eq!(chain, ChainType::Forward);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reload_rules_clears_outbound_allow_records() {
|
||||
let filter = AclFilter::new();
|
||||
filter.outbound_allow_records.insert(
|
||||
OutboundAllowRecord {
|
||||
src_ip: IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)),
|
||||
dst_ip: IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2)),
|
||||
src_port: Some(1234),
|
||||
dst_port: Some(80),
|
||||
protocol: Protocol::Tcp,
|
||||
},
|
||||
Instant::now(),
|
||||
);
|
||||
assert_eq!(filter.outbound_allow_records.len(), 1);
|
||||
|
||||
filter.reload_rules(Some(&Acl::default()));
|
||||
|
||||
assert_eq!(filter.outbound_allow_records.len(), 0);
|
||||
|
||||
filter.outbound_allow_records.insert(
|
||||
OutboundAllowRecord {
|
||||
src_ip: IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2)),
|
||||
dst_ip: IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1)),
|
||||
src_port: Some(4321),
|
||||
dst_port: Some(443),
|
||||
protocol: Protocol::Tcp,
|
||||
},
|
||||
Instant::now(),
|
||||
);
|
||||
assert_eq!(filter.outbound_allow_records.len(), 1);
|
||||
|
||||
filter.reload_rules(None);
|
||||
|
||||
assert_eq!(filter.outbound_allow_records.len(), 0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
//! Access-control list packet filtering: the per-rule processor and the
|
||||
//! filter wiring it into the peer/NIC packet pipelines.
|
||||
|
||||
pub(crate) mod filter;
|
||||
pub(crate) mod processor;
|
||||
|
||||
pub(crate) use filter::AclFilter;
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,131 @@
|
||||
use std::sync::{Arc, Weak};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::{
|
||||
connectivity::protocol::raw,
|
||||
events::{CoreEvent, CoreEventSink},
|
||||
listener::{
|
||||
AcceptedSocketHandler,
|
||||
transport::{AcceptedTransport, AcceptedTunnelHandler},
|
||||
},
|
||||
tunnel::Tunnel,
|
||||
};
|
||||
|
||||
use super::peer_manager::PeerManagerCore;
|
||||
|
||||
pub(crate) struct PeerAcceptedTunnelHandler {
|
||||
peer_manager: Weak<PeerManagerCore>,
|
||||
events: Arc<dyn CoreEventSink>,
|
||||
}
|
||||
|
||||
impl PeerAcceptedTunnelHandler {
|
||||
pub(crate) fn new(
|
||||
peer_manager: &Arc<PeerManagerCore>,
|
||||
events: Arc<dyn CoreEventSink>,
|
||||
) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
peer_manager: Arc::downgrade(peer_manager),
|
||||
events,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AcceptedTunnelHandler for PeerAcceptedTunnelHandler {
|
||||
async fn handle_tunnel(&self, tunnel: Box<dyn Tunnel>) -> anyhow::Result<()> {
|
||||
let tunnel_info = tunnel
|
||||
.info()
|
||||
.ok_or_else(|| anyhow::anyhow!("accepted tunnel has no tunnel info"))?;
|
||||
let local_url = tunnel_info
|
||||
.local_addr
|
||||
.clone()
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
let remote_url = tunnel_info
|
||||
.remote_addr
|
||||
.clone()
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
self.events.emit(CoreEvent::TunnelAccepted {
|
||||
local_url: local_url.clone(),
|
||||
remote_url: remote_url.clone(),
|
||||
});
|
||||
tracing::info!(ret = ?tunnel, "conn accepted");
|
||||
|
||||
let Some(peer_manager) = self.peer_manager.upgrade() else {
|
||||
let error = "peer manager is gone, cannot handle tunnel".to_owned();
|
||||
self.events.emit(CoreEvent::TunnelAdmissionFailed {
|
||||
local_url,
|
||||
remote_url,
|
||||
error: error.clone(),
|
||||
});
|
||||
tracing::error!(error = %error, "handle conn error");
|
||||
return Err(anyhow::anyhow!(error));
|
||||
};
|
||||
if let Err(error) = peer_manager.add_tunnel_as_server(tunnel, true).await {
|
||||
self.events.emit(CoreEvent::TunnelAdmissionFailed {
|
||||
local_url,
|
||||
remote_url,
|
||||
error: error.to_string(),
|
||||
});
|
||||
tracing::error!(?error, "handle conn error");
|
||||
return Err(error.into());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) struct RawAcceptedTransportHandler {
|
||||
peer_manager: Weak<PeerManagerCore>,
|
||||
}
|
||||
|
||||
impl RawAcceptedTransportHandler {
|
||||
pub(crate) fn new(peer_manager: &Arc<PeerManagerCore>) -> Self {
|
||||
Self {
|
||||
peer_manager: Arc::downgrade(peer_manager),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl<TcpSocket> AcceptedSocketHandler<AcceptedTransport<TcpSocket>> for RawAcceptedTransportHandler
|
||||
where
|
||||
TcpSocket: crate::socket::tcp::VirtualTcpSocket,
|
||||
{
|
||||
async fn handle_accepted_socket(
|
||||
&self,
|
||||
accepted: AcceptedTransport<TcpSocket>,
|
||||
) -> anyhow::Result<()> {
|
||||
let peer_manager = self
|
||||
.peer_manager
|
||||
.upgrade()
|
||||
.ok_or_else(|| anyhow::anyhow!("peer manager is gone"))?;
|
||||
let tunnel = match accepted {
|
||||
AcceptedTransport::Tunnel { tunnel, .. } => tunnel,
|
||||
AcceptedTransport::Tcp {
|
||||
socket, local_url, ..
|
||||
} => {
|
||||
if local_url.scheme() != "tcp" {
|
||||
anyhow::bail!("unsupported raw TCP listener protocol: {local_url}");
|
||||
}
|
||||
raw::upgrade_accepted_tcp_with_local_url(socket, local_url)?
|
||||
}
|
||||
AcceptedTransport::Udp {
|
||||
session, local_url, ..
|
||||
} => {
|
||||
if local_url.scheme() != "udp" {
|
||||
anyhow::bail!("unsupported raw UDP listener protocol: {local_url}");
|
||||
}
|
||||
raw::upgrade_accepted_udp_with_local_url(session, local_url)?
|
||||
}
|
||||
AcceptedTransport::ByteStream {
|
||||
socket,
|
||||
local_url,
|
||||
remote_url,
|
||||
} => raw::upgrade_accepted_byte_stream(socket, local_url, remote_url)?,
|
||||
};
|
||||
peer_manager.add_tunnel_as_server(tunnel, true).await?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
//! Peer connection primitives: noise sessions, individual peer connections,
|
||||
//! and the peer map that multiplexes them.
|
||||
|
||||
pub(crate) mod peer;
|
||||
pub(crate) mod peer_conn;
|
||||
pub(crate) mod peer_conn_ping;
|
||||
pub(crate) mod peer_map;
|
||||
pub(crate) mod peer_session;
|
||||
@@ -0,0 +1,290 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use crossbeam::atomic::AtomicCell;
|
||||
use dashmap::{DashMap, DashSet};
|
||||
use parking_lot::RwLock;
|
||||
|
||||
use tokio::{select, sync::mpsc};
|
||||
|
||||
use tracing::Instrument;
|
||||
|
||||
use super::peer_conn::{PeerConn, PeerConnId};
|
||||
use crate::peers::{
|
||||
PacketRecvChan,
|
||||
context::{ArcPeerContext, PeerEvent},
|
||||
util::shrink_dashmap,
|
||||
};
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
packet::ZCPacket,
|
||||
peers::error::Error,
|
||||
proto::{core_peer::peer::PeerConnInfo, peer_rpc::PeerIdentityType},
|
||||
};
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
|
||||
type ArcPeerConn = Arc<PeerConn>;
|
||||
type ConnMap = Arc<DashMap<PeerConnId, ArcPeerConn>>;
|
||||
|
||||
pub struct Peer {
|
||||
pub peer_node_id: PeerId,
|
||||
conns: ConnMap,
|
||||
context: ArcPeerContext,
|
||||
|
||||
packet_recv_chan: PacketRecvChan,
|
||||
|
||||
close_event_sender: mpsc::Sender<PeerConnId>,
|
||||
#[allow(dead_code)]
|
||||
close_event_listener: AbortOnDropHandle<()>,
|
||||
|
||||
shutdown_notifier: Arc<tokio::sync::Notify>,
|
||||
|
||||
default_conn_id: Arc<AtomicCell<PeerConnId>>,
|
||||
peer_identity_type: Arc<AtomicCell<Option<PeerIdentityType>>>,
|
||||
peer_public_key: Arc<RwLock<Option<Vec<u8>>>>,
|
||||
#[allow(dead_code)]
|
||||
default_conn_id_clear_task: AbortOnDropHandle<()>,
|
||||
}
|
||||
|
||||
impl Peer {
|
||||
pub(crate) fn new(
|
||||
peer_node_id: PeerId,
|
||||
packet_recv_chan: PacketRecvChan,
|
||||
context: ArcPeerContext,
|
||||
) -> Self {
|
||||
let conns: ConnMap = Arc::new(DashMap::new());
|
||||
let (close_event_sender, mut close_event_receiver) = mpsc::channel(10);
|
||||
let shutdown_notifier = Arc::new(tokio::sync::Notify::new());
|
||||
let peer_identity_type = Arc::new(AtomicCell::new(None));
|
||||
let peer_identity_type_copy = peer_identity_type.clone();
|
||||
let peer_public_key = Arc::new(RwLock::new(None));
|
||||
let peer_public_key_copy = peer_public_key.clone();
|
||||
|
||||
let conns_copy = conns.clone();
|
||||
let shutdown_notifier_copy = shutdown_notifier.clone();
|
||||
let context_copy = context.clone();
|
||||
let close_event_listener = AbortOnDropHandle::new(tokio::spawn(
|
||||
async move {
|
||||
loop {
|
||||
select! {
|
||||
ret = close_event_receiver.recv() => {
|
||||
if ret.is_none() {
|
||||
break;
|
||||
}
|
||||
let ret = ret.unwrap();
|
||||
tracing::warn!(
|
||||
?peer_node_id,
|
||||
?ret,
|
||||
"notified that peer conn is closed",
|
||||
);
|
||||
|
||||
if let Some((_, conn)) = conns_copy.remove(&ret) {
|
||||
context_copy.issue_event(PeerEvent::PeerConnRemoved(
|
||||
conn.get_conn_info(),
|
||||
));
|
||||
shrink_dashmap(&conns_copy, Some(4));
|
||||
if conns_copy.is_empty() {
|
||||
peer_identity_type_copy.store(None);
|
||||
*peer_public_key_copy.write() = None;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
_ = shutdown_notifier_copy.notified() => {
|
||||
close_event_receiver.close();
|
||||
tracing::warn!(?peer_node_id, "peer close event listener notified");
|
||||
}
|
||||
}
|
||||
}
|
||||
tracing::info!("peer {} close event listener exit", peer_node_id);
|
||||
}
|
||||
.instrument(tracing::info_span!(
|
||||
"peer_close_event_listener",
|
||||
?peer_node_id,
|
||||
)),
|
||||
));
|
||||
|
||||
let default_conn_id = Arc::new(AtomicCell::new(PeerConnId::default()));
|
||||
|
||||
let conns_copy = conns.clone();
|
||||
let default_conn_id_copy = default_conn_id.clone();
|
||||
let default_conn_id_clear_task = AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
loop {
|
||||
crate::foundation::time::sleep(std::time::Duration::from_secs(5)).await;
|
||||
if conns_copy.len() > 1 {
|
||||
default_conn_id_copy.store(PeerConnId::default());
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
Peer {
|
||||
peer_node_id,
|
||||
conns,
|
||||
packet_recv_chan,
|
||||
context,
|
||||
|
||||
close_event_sender,
|
||||
close_event_listener,
|
||||
|
||||
shutdown_notifier,
|
||||
default_conn_id,
|
||||
peer_identity_type,
|
||||
peer_public_key,
|
||||
default_conn_id_clear_task,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn add_peer_conn(&self, mut conn: PeerConn) -> Result<(), Error> {
|
||||
let conn_identity_type = conn.get_peer_identity_type();
|
||||
let peer_identity_type = self.peer_identity_type.load();
|
||||
if let Some(peer_identity_type) = peer_identity_type {
|
||||
if peer_identity_type != conn_identity_type {
|
||||
return Err(Error::SecretKeyError(format!(
|
||||
"peer identity type mismatch. peer: {:?}, conn: {:?}",
|
||||
peer_identity_type, conn_identity_type
|
||||
)));
|
||||
}
|
||||
} else {
|
||||
self.peer_identity_type.store(Some(conn_identity_type));
|
||||
}
|
||||
|
||||
let close_notifier = conn.get_close_notifier();
|
||||
let conn_info = conn.get_conn_info();
|
||||
let conn_pubkey = conn_info.noise_remote_static_pubkey.clone();
|
||||
{
|
||||
let mut peer_pubkey = self.peer_public_key.write();
|
||||
if let Some(existing_pubkey) = peer_pubkey.as_ref() {
|
||||
if existing_pubkey != &conn_pubkey {
|
||||
return Err(Error::SecretKeyError(format!(
|
||||
"peer public key mismatch. peer_id: {}, existing_len: {}, new_len: {}",
|
||||
self.peer_node_id,
|
||||
existing_pubkey.len(),
|
||||
conn_pubkey.len()
|
||||
)));
|
||||
}
|
||||
} else {
|
||||
*peer_pubkey = Some(conn_pubkey);
|
||||
}
|
||||
}
|
||||
|
||||
conn.start_recv_loop(self.packet_recv_chan.clone()).await;
|
||||
conn.start_pingpong();
|
||||
self.conns.insert(conn.get_conn_id(), Arc::new(conn));
|
||||
|
||||
let close_event_sender = self.close_event_sender.clone();
|
||||
tokio::spawn(async move {
|
||||
let conn_id = close_notifier.get_conn_id();
|
||||
if let Some(mut waiter) = close_notifier.get_waiter().await {
|
||||
let _ = waiter.recv().await;
|
||||
}
|
||||
if let Err(e) = close_event_sender.send(conn_id).await {
|
||||
tracing::warn!(?conn_id, "failed to send close event: {}", e);
|
||||
}
|
||||
});
|
||||
|
||||
self.context
|
||||
.issue_event(PeerEvent::PeerConnAdded(conn_info));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn select_conn(&self) -> Option<ArcPeerConn> {
|
||||
let default_conn_id = self.default_conn_id.load();
|
||||
if let Some(conn) = self.conns.get(&default_conn_id) {
|
||||
return Some(conn.clone());
|
||||
}
|
||||
|
||||
// find a conn with the smallest latency
|
||||
let mut min_latency = u64::MAX;
|
||||
for conn in self.conns.iter() {
|
||||
let latency = conn.value().get_stats().latency_us;
|
||||
if latency < min_latency {
|
||||
min_latency = latency;
|
||||
self.default_conn_id.store(conn.get_conn_id());
|
||||
}
|
||||
}
|
||||
|
||||
self.conns
|
||||
.get(&self.default_conn_id.load())
|
||||
.map(|conn| conn.clone())
|
||||
}
|
||||
|
||||
pub async fn send_msg(&self, msg: ZCPacket) -> Result<(), Error> {
|
||||
let Some(conn) = self.select_conn().await else {
|
||||
return Err(Error::PeerNoConnectionError(self.peer_node_id));
|
||||
};
|
||||
conn.send_msg(msg).await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn close_peer_conn(&self, conn_id: &PeerConnId) -> Result<(), Error> {
|
||||
let has_key = self.conns.contains_key(conn_id);
|
||||
if !has_key {
|
||||
return Err(Error::NotFound);
|
||||
}
|
||||
self.close_event_sender.send(*conn_id).await.unwrap();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn list_peer_conns(&self) -> Vec<PeerConnInfo> {
|
||||
let mut conns = vec![];
|
||||
for conn in self.conns.iter() {
|
||||
// do not lock here, otherwise it will cause dashmap deadlock
|
||||
conns.push(conn.clone());
|
||||
}
|
||||
|
||||
let mut ret = Vec::new();
|
||||
for conn in conns {
|
||||
let info = conn.get_conn_info();
|
||||
if !info.is_closed {
|
||||
ret.push(info);
|
||||
} else {
|
||||
let conn_id = info.conn_id.parse().unwrap();
|
||||
let _ = self.close_peer_conn(&conn_id).await;
|
||||
}
|
||||
}
|
||||
ret
|
||||
}
|
||||
|
||||
pub fn has_live_conns(&self) -> bool {
|
||||
self.conns.iter().any(|entry| !entry.value().is_closed())
|
||||
}
|
||||
|
||||
pub fn has_directly_connected_conn(&self) -> bool {
|
||||
self.conns
|
||||
.iter()
|
||||
.any(|entry| !entry.value().is_closed() && !entry.value().is_hole_punched())
|
||||
}
|
||||
|
||||
pub fn get_directly_connections(&self) -> DashSet<uuid::Uuid> {
|
||||
self.conns
|
||||
.iter()
|
||||
.filter(|entry| !(entry.value()).is_hole_punched())
|
||||
.map(|entry| (entry.value()).get_conn_id())
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn get_default_conn_id(&self) -> PeerConnId {
|
||||
self.default_conn_id.load()
|
||||
}
|
||||
|
||||
pub fn get_peer_identity_type(&self) -> Option<PeerIdentityType> {
|
||||
self.peer_identity_type.load()
|
||||
}
|
||||
|
||||
pub fn get_peer_public_key(&self) -> Option<Vec<u8>> {
|
||||
self.peer_public_key.read().clone()
|
||||
}
|
||||
}
|
||||
|
||||
// pritn on drop
|
||||
impl Drop for Peer {
|
||||
fn drop(&mut self) {
|
||||
self.conns.retain(|_, conn| {
|
||||
self.context
|
||||
.issue_event(PeerEvent::PeerConnRemoved(conn.get_conn_info()));
|
||||
false
|
||||
});
|
||||
self.shutdown_notifier.notify_one();
|
||||
tracing::info!("peer {} drop", self.peer_node_id);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,364 @@
|
||||
use std::{
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicU32, Ordering},
|
||||
},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use rand::{Rng, thread_rng};
|
||||
use tokio::{sync::broadcast, task::JoinSet};
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
foundation::time::{Interval, interval, timeout},
|
||||
packet::{PacketType, ZCPacket},
|
||||
peers::{context::ArcPeerContext, error::Error},
|
||||
tunnel::{
|
||||
TunnelError,
|
||||
mpsc::MpscTunnelSender,
|
||||
stats::{Throughput, WindowLatency},
|
||||
},
|
||||
};
|
||||
|
||||
struct PingIntervalController {
|
||||
throughput: Arc<Throughput>,
|
||||
loss_counter: Arc<AtomicU32>,
|
||||
|
||||
interval: Interval,
|
||||
|
||||
logic_time: u64,
|
||||
last_send_logic_time: u64,
|
||||
|
||||
backoff_idx: i32,
|
||||
max_backoff_idx: i32,
|
||||
|
||||
last_throughput: Throughput,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PingIntervalController {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("PingIntervalController")
|
||||
.field("throughput", &self.throughput)
|
||||
.field("loss_counter", &self.loss_counter)
|
||||
.field("logic_time", &self.logic_time)
|
||||
.field("last_send_logic_time", &self.last_send_logic_time)
|
||||
.field("backoff_idx", &self.backoff_idx)
|
||||
.field("max_backoff_idx", &self.max_backoff_idx)
|
||||
.field("last_throughput", &self.last_throughput)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl PingIntervalController {
|
||||
fn new(throughput: Arc<Throughput>, loss_counter: Arc<AtomicU32>) -> Self {
|
||||
let last_throughput = (*throughput).clone();
|
||||
|
||||
Self {
|
||||
throughput,
|
||||
loss_counter,
|
||||
interval: interval(Duration::from_secs(1)),
|
||||
logic_time: 0,
|
||||
last_send_logic_time: 0,
|
||||
|
||||
backoff_idx: 0,
|
||||
max_backoff_idx: 5,
|
||||
|
||||
last_throughput,
|
||||
}
|
||||
}
|
||||
|
||||
async fn tick(&mut self) {
|
||||
self.interval.tick().await;
|
||||
self.logic_time += 1;
|
||||
}
|
||||
|
||||
fn tx_increase(&self) -> bool {
|
||||
self.throughput.tx_packets() > self.last_throughput.tx_packets()
|
||||
}
|
||||
|
||||
fn rx_increase(&self) -> bool {
|
||||
self.throughput.rx_packets() > self.last_throughput.rx_packets()
|
||||
}
|
||||
|
||||
fn should_send_ping(&mut self) -> bool {
|
||||
if self.loss_counter.load(Ordering::Relaxed) > 0 {
|
||||
self.backoff_idx = 0;
|
||||
} else if self.tx_increase() && !self.rx_increase() {
|
||||
// if tx increase but rx not increase, we should do pingpong more frequently
|
||||
self.backoff_idx = 0;
|
||||
}
|
||||
|
||||
self.last_throughput = (*self.throughput).clone();
|
||||
|
||||
if (self.logic_time - self.last_send_logic_time) < (1 << self.backoff_idx) {
|
||||
return false;
|
||||
}
|
||||
|
||||
self.backoff_idx = std::cmp::min(self.backoff_idx + 1, self.max_backoff_idx);
|
||||
|
||||
// use this makes two peers not pingpong at the same time
|
||||
if self.backoff_idx > self.max_backoff_idx - 2 && thread_rng().gen_bool(0.2) {
|
||||
self.backoff_idx -= 1;
|
||||
}
|
||||
|
||||
self.last_send_logic_time = self.logic_time;
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PeerConnPinger {
|
||||
my_peer_id: PeerId,
|
||||
peer_id: PeerId,
|
||||
sink: MpscTunnelSender,
|
||||
ctrl_sender: broadcast::Sender<ZCPacket>,
|
||||
latency_stats: Arc<WindowLatency>,
|
||||
loss_rate_stats: Arc<AtomicU32>,
|
||||
throughput_stats: Arc<Throughput>,
|
||||
context: ArcPeerContext,
|
||||
network_name: String,
|
||||
tasks: JoinSet<Result<(), TunnelError>>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PeerConnPinger {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("PeerConnPinger")
|
||||
.field("my_peer_id", &self.my_peer_id)
|
||||
.field("peer_id", &self.peer_id)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerConnPinger {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) fn new(
|
||||
my_peer_id: PeerId,
|
||||
peer_id: PeerId,
|
||||
sink: MpscTunnelSender,
|
||||
ctrl_sender: broadcast::Sender<ZCPacket>,
|
||||
latency_stats: Arc<WindowLatency>,
|
||||
loss_rate_stats: Arc<AtomicU32>,
|
||||
throughput_stats: Arc<Throughput>,
|
||||
context: ArcPeerContext,
|
||||
network_name: String,
|
||||
) -> Self {
|
||||
Self {
|
||||
my_peer_id,
|
||||
peer_id,
|
||||
sink,
|
||||
tasks: JoinSet::new(),
|
||||
latency_stats,
|
||||
ctrl_sender,
|
||||
loss_rate_stats,
|
||||
throughput_stats,
|
||||
context,
|
||||
network_name,
|
||||
}
|
||||
}
|
||||
|
||||
fn new_ping_packet(my_node_id: PeerId, peer_id: PeerId, seq: u32) -> ZCPacket {
|
||||
let mut packet = ZCPacket::new_with_payload(&seq.to_le_bytes());
|
||||
packet.fill_peer_manager_hdr(my_node_id, peer_id, PacketType::Ping as u8);
|
||||
packet
|
||||
}
|
||||
|
||||
async fn do_pingpong_once(
|
||||
my_node_id: PeerId,
|
||||
peer_id: PeerId,
|
||||
sink: &MpscTunnelSender,
|
||||
context: &ArcPeerContext,
|
||||
network_name: &str,
|
||||
receiver: &mut broadcast::Receiver<ZCPacket>,
|
||||
seq: u32,
|
||||
) -> Result<u128, Error> {
|
||||
// should add seq here. so latency can be calculated more accurately
|
||||
let req = Self::new_ping_packet(my_node_id, peer_id, seq);
|
||||
let req_len = req.buf_len() as u64;
|
||||
sink.send(req).await?;
|
||||
context.record_control_tx(network_name, req_len);
|
||||
|
||||
let now = Instant::now();
|
||||
// wait until we get a pong packet in ctrl_resp_receiver
|
||||
let resp = timeout(Duration::from_secs(2), async {
|
||||
loop {
|
||||
match receiver.recv().await {
|
||||
Ok(p) => {
|
||||
let payload = p.payload();
|
||||
let Ok(seq_buf) = payload[0..4].try_into() else {
|
||||
tracing::debug!("pingpong recv invalid packet, continue");
|
||||
continue;
|
||||
};
|
||||
let resp_seq = u32::from_le_bytes(seq_buf);
|
||||
if resp_seq == seq {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
return Err(Error::WaitRespError(format!(
|
||||
"wait ping response error: {:?}",
|
||||
e
|
||||
)));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
})
|
||||
.await;
|
||||
|
||||
tracing::trace!(?resp, "wait ping response done");
|
||||
|
||||
if resp.is_err() {
|
||||
return Err(Error::WaitRespError(
|
||||
"wait ping response timeout".to_owned(),
|
||||
));
|
||||
}
|
||||
|
||||
if resp.as_ref().unwrap().is_err() {
|
||||
return Err(resp.unwrap().err().unwrap());
|
||||
}
|
||||
|
||||
Ok(now.elapsed().as_micros())
|
||||
}
|
||||
|
||||
pub async fn pingpong(&mut self) {
|
||||
let sink = self.sink.clone();
|
||||
let context = self.context.clone();
|
||||
let network_name = self.network_name.clone();
|
||||
let my_node_id = self.my_peer_id;
|
||||
let peer_id = self.peer_id;
|
||||
let latency_stats = self.latency_stats.clone();
|
||||
|
||||
let (ping_res_sender, mut ping_res_receiver) = tokio::sync::mpsc::channel(100);
|
||||
|
||||
// one with 1% precision
|
||||
let loss_rate_stats_1 = WindowLatency::new(100);
|
||||
// disconnect the connection if lost 5 pingpong consecutively
|
||||
let loss_counter = Arc::new(AtomicU32::new(0));
|
||||
|
||||
let stopped = Arc::new(AtomicU32::new(0));
|
||||
|
||||
// generate a pingpong task every 200ms
|
||||
let mut pingpong_tasks = JoinSet::new();
|
||||
let ctrl_resp_sender = self.ctrl_sender.clone();
|
||||
let stopped_clone = stopped.clone();
|
||||
let mut controller =
|
||||
PingIntervalController::new(self.throughput_stats.clone(), loss_counter.clone());
|
||||
self.tasks.spawn(
|
||||
async move {
|
||||
let mut req_seq = 0;
|
||||
loop {
|
||||
controller.tick().await;
|
||||
|
||||
if stopped_clone.load(Ordering::Relaxed) != 0 {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
while pingpong_tasks.len() > 5 {
|
||||
pingpong_tasks.join_next().await;
|
||||
}
|
||||
|
||||
if !controller.should_send_ping() {
|
||||
continue;
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
"pingpong controller send pingpong task, seq: {}, node_id: {}, controller: {:?}",
|
||||
req_seq,
|
||||
my_node_id,
|
||||
controller,
|
||||
);
|
||||
|
||||
let sink = sink.clone();
|
||||
let context = context.clone();
|
||||
let network_name = network_name.clone();
|
||||
let receiver = ctrl_resp_sender.subscribe();
|
||||
let ping_res_sender = ping_res_sender.clone();
|
||||
pingpong_tasks.spawn(async move {
|
||||
let mut receiver = receiver.resubscribe();
|
||||
let pingpong_once_ret = Self::do_pingpong_once(
|
||||
my_node_id,
|
||||
peer_id,
|
||||
&sink,
|
||||
&context,
|
||||
&network_name,
|
||||
&mut receiver,
|
||||
req_seq,
|
||||
)
|
||||
.await;
|
||||
|
||||
if let Err(e) = ping_res_sender.send(pingpong_once_ret).await {
|
||||
tracing::info!(?e, "pingpong task send result error, exit..");
|
||||
};
|
||||
});
|
||||
|
||||
req_seq = req_seq.wrapping_add(1);
|
||||
}
|
||||
}
|
||||
.instrument(tracing::info_span!(
|
||||
"pingpong_controller",
|
||||
?my_node_id,
|
||||
?peer_id
|
||||
)),
|
||||
);
|
||||
|
||||
let throughput = self.throughput_stats.clone();
|
||||
let mut last_rx_packets = throughput.rx_packets();
|
||||
|
||||
while let Some(ret) = ping_res_receiver.recv().await {
|
||||
if let Ok(lat) = ret {
|
||||
latency_stats.record_latency(lat as u32);
|
||||
|
||||
loss_rate_stats_1.record_latency(0);
|
||||
} else {
|
||||
loss_rate_stats_1.record_latency(1);
|
||||
loss_counter.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
let loss_rate_1: f64 = loss_rate_stats_1.get_latency_us();
|
||||
|
||||
tracing::trace!(
|
||||
?ret,
|
||||
?self,
|
||||
?loss_rate_1,
|
||||
"pingpong task recv pingpong_once result"
|
||||
);
|
||||
|
||||
let current_rx_packets = throughput.rx_packets();
|
||||
if last_rx_packets != current_rx_packets {
|
||||
// if we receive some packet from peers, reset the counter to avoid conn close.
|
||||
// conn will close only if we have 5 continous round pingpong loss after no packet received.
|
||||
loss_counter.store(0, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
"loss_counter: {:?}, loss_rate_1: {}, cur_rx_packets: {}, last_rx: {}, node_id: {}",
|
||||
loss_counter,
|
||||
loss_rate_1,
|
||||
current_rx_packets,
|
||||
last_rx_packets,
|
||||
my_node_id
|
||||
);
|
||||
|
||||
if loss_counter.load(Ordering::Relaxed) >= 5 {
|
||||
tracing::warn!(
|
||||
?ret,
|
||||
?self,
|
||||
?loss_rate_1,
|
||||
?loss_counter,
|
||||
?last_rx_packets,
|
||||
?current_rx_packets,
|
||||
"pingpong loss too much pingpong packet and no other ingress packets, closing the connection",
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
last_rx_packets = throughput.rx_packets();
|
||||
self.loss_rate_stats
|
||||
.store((loss_rate_1 * 100.0) as u32, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
stopped.store(1, Ordering::Relaxed);
|
||||
ping_res_receiver.close();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,463 @@
|
||||
use std::{
|
||||
collections::{BTreeSet, HashMap, HashSet},
|
||||
net::{Ipv4Addr, Ipv6Addr},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use anyhow::Context;
|
||||
use dashmap::{DashMap, DashSet};
|
||||
use parking_lot::Mutex;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
packet::ZCPacket,
|
||||
peers::{
|
||||
context::{ArcPeerContext, NetworkIdentity, PeerEvent},
|
||||
error::Error,
|
||||
util::shrink_dashmap,
|
||||
},
|
||||
proto::{
|
||||
core_peer::peer::{PeerConnInfo, Route as CoreRoute},
|
||||
peer_rpc::{
|
||||
DirectConnectedPeerInfo, PeerIdentityType, PeerInfoForGlobalMap, RoutePeerInfo,
|
||||
},
|
||||
},
|
||||
tunnel::TunnelError,
|
||||
};
|
||||
|
||||
use super::{
|
||||
peer::Peer,
|
||||
peer_conn::{PeerConn, PeerConnId},
|
||||
};
|
||||
use crate::peers::{
|
||||
PacketRecvChan,
|
||||
route::{ArcRoute, NextHopPolicy},
|
||||
};
|
||||
|
||||
pub struct PeerMap {
|
||||
context: ArcPeerContext,
|
||||
my_peer_id: PeerId,
|
||||
peer_map: DashMap<PeerId, Arc<Peer>>,
|
||||
packet_send: PacketRecvChan,
|
||||
routes: RwLock<Vec<ArcRoute>>,
|
||||
alive_client_urls: Arc<Mutex<HashMap<url::Url, HashSet<PeerConnId>>>>,
|
||||
}
|
||||
|
||||
impl PeerMap {
|
||||
pub(crate) fn new(
|
||||
packet_send: PacketRecvChan,
|
||||
context: ArcPeerContext,
|
||||
my_peer_id: PeerId,
|
||||
) -> Self {
|
||||
PeerMap {
|
||||
context,
|
||||
my_peer_id,
|
||||
peer_map: DashMap::new(),
|
||||
packet_send,
|
||||
routes: RwLock::new(Vec::new()),
|
||||
alive_client_urls: Arc::new(Mutex::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn add_new_peer(&self, peer: Peer) {
|
||||
let peer_id = peer.peer_node_id;
|
||||
self.peer_map.insert(peer_id, Arc::new(peer));
|
||||
self.context.issue_event(PeerEvent::PeerAdded(peer_id));
|
||||
}
|
||||
|
||||
pub async fn add_new_peer_conn(&self, peer_conn: PeerConn) -> Result<(), Error> {
|
||||
let _ = self.maintain_alive_client_urls(&peer_conn);
|
||||
let peer_id = peer_conn.get_peer_id();
|
||||
let no_entry = self.peer_map.get(&peer_id).is_none();
|
||||
if no_entry {
|
||||
let new_peer = Peer::new(peer_id, self.packet_send.clone(), self.context.clone());
|
||||
new_peer.add_peer_conn(peer_conn).await?;
|
||||
self.add_new_peer(new_peer).await;
|
||||
} else {
|
||||
let peer = self.peer_map.get(&peer_id).unwrap().clone();
|
||||
peer.add_peer_conn(peer_conn).await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn maintain_alive_client_urls(&self, peer_conn: &PeerConn) -> Option<()> {
|
||||
let conn_info = peer_conn.get_conn_info();
|
||||
if !conn_info.is_client {
|
||||
return None;
|
||||
}
|
||||
|
||||
let close_notifier = peer_conn.get_close_notifier();
|
||||
let alive_conns_weak = Arc::downgrade(&self.alive_client_urls);
|
||||
let conn_id = close_notifier.get_conn_id();
|
||||
let alive_client_url: url::Url = conn_info.tunnel?.remote_addr?.into();
|
||||
self.alive_client_urls
|
||||
.lock()
|
||||
.entry(alive_client_url.clone())
|
||||
.or_default()
|
||||
.insert(conn_id);
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Some(mut waiter) = close_notifier.get_waiter().await {
|
||||
let _ = waiter.recv().await;
|
||||
}
|
||||
let Some(alive_conns) = alive_conns_weak.upgrade() else {
|
||||
return;
|
||||
};
|
||||
let mut guard = alive_conns.lock();
|
||||
if let Some(conn_ids) = guard.get_mut(&alive_client_url) {
|
||||
conn_ids.retain(|id| id != &conn_id);
|
||||
if conn_ids.is_empty() {
|
||||
guard.remove(&alive_client_url);
|
||||
}
|
||||
}
|
||||
let alive_conn_count = guard.len();
|
||||
drop(guard);
|
||||
tracing::debug!(
|
||||
?conn_id,
|
||||
"peer conn is closed, current alive conns: {}",
|
||||
alive_conn_count
|
||||
);
|
||||
});
|
||||
|
||||
Some(())
|
||||
}
|
||||
|
||||
pub fn is_client_url_alive(&self, url: &url::Url) -> bool {
|
||||
self.alive_client_urls.lock().contains_key(url)
|
||||
}
|
||||
|
||||
pub fn get_peer_by_id(&self, peer_id: PeerId) -> Option<Arc<Peer>> {
|
||||
self.peer_map.get(&peer_id).map(|v| v.clone())
|
||||
}
|
||||
|
||||
pub fn get_directly_connections_by_peer_id(&self, peer_id: PeerId) -> DashSet<uuid::Uuid> {
|
||||
if let Some(peer) = self.get_peer_by_id(peer_id) {
|
||||
return peer.get_directly_connections();
|
||||
}
|
||||
|
||||
DashSet::new()
|
||||
}
|
||||
|
||||
pub fn has_peer(&self, peer_id: PeerId) -> bool {
|
||||
peer_id == self.my_peer_id || self.peer_map.contains_key(&peer_id)
|
||||
}
|
||||
|
||||
pub async fn send_msg_directly(&self, msg: ZCPacket, dst_peer_id: PeerId) -> Result<(), Error> {
|
||||
if dst_peer_id == self.my_peer_id {
|
||||
let packet_send = self.packet_send.clone();
|
||||
tokio::spawn(async move {
|
||||
let ret = packet_send
|
||||
.send(msg)
|
||||
.await
|
||||
.with_context(|| "send msg to self failed");
|
||||
if ret.is_err() {
|
||||
tracing::error!("send msg to self failed: {:?}", ret);
|
||||
}
|
||||
});
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
match self.get_peer_by_id(dst_peer_id) {
|
||||
Some(peer) => {
|
||||
peer.send_msg(msg).await?;
|
||||
}
|
||||
None => {
|
||||
tracing::error!("no peer for dst_peer_id: {}", dst_peer_id);
|
||||
return Err(Error::RouteError(Some(format!(
|
||||
"peer map sengmsg directly no connected dst_peer_id: {}",
|
||||
dst_peer_id
|
||||
))));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_gateway_peer_id(
|
||||
&self,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Option<PeerId> {
|
||||
if dst_peer_id == self.my_peer_id {
|
||||
return Some(dst_peer_id);
|
||||
}
|
||||
|
||||
if self.has_peer(dst_peer_id) && matches!(policy, NextHopPolicy::LeastHop) {
|
||||
return Some(dst_peer_id);
|
||||
}
|
||||
|
||||
// get route info
|
||||
for route in self.routes.read().await.iter() {
|
||||
if let Some(gateway_peer_id) = route
|
||||
.get_next_hop_with_policy(dst_peer_id, policy.clone())
|
||||
.await
|
||||
{
|
||||
// NOTIC: for foreign network, gateway_peer_id may not connect to me
|
||||
return Some(gateway_peer_id);
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
pub async fn list_peers_own_foreign_network(
|
||||
&self,
|
||||
network_identity: &NetworkIdentity,
|
||||
) -> Vec<PeerId> {
|
||||
let mut ret = Vec::new();
|
||||
for route in self.routes.read().await.iter() {
|
||||
let peers = route.list_peers_own_foreign_network(network_identity).await;
|
||||
ret.extend(peers);
|
||||
}
|
||||
ret
|
||||
}
|
||||
|
||||
pub async fn send_msg(
|
||||
&self,
|
||||
msg: ZCPacket,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<(), Error> {
|
||||
let Some(gateway_peer_id) = self.get_gateway_peer_id(dst_peer_id, policy).await else {
|
||||
return Err(Error::RouteError(Some(format!(
|
||||
"peer map sengmsg no gateway for dst_peer_id: {}",
|
||||
dst_peer_id
|
||||
))));
|
||||
};
|
||||
|
||||
self.send_msg_directly(msg, gateway_peer_id).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_peer_id_by_ipv4(&self, ipv4: &Ipv4Addr) -> Option<PeerId> {
|
||||
for route in self.routes.read().await.iter() {
|
||||
let peer_id = route.get_peer_id_by_ipv4(ipv4).await;
|
||||
if peer_id.is_some() {
|
||||
return peer_id;
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub async fn get_peer_id_by_ipv6(&self, ipv6: &Ipv6Addr) -> Option<PeerId> {
|
||||
for route in self.routes.read().await.iter() {
|
||||
let peer_id = route.get_peer_id_by_ipv6(ipv6).await;
|
||||
if peer_id.is_some() {
|
||||
return peer_id;
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub async fn get_route_peer_info(&self, peer_id: PeerId) -> Option<RoutePeerInfo> {
|
||||
for route in self.routes.read().await.iter() {
|
||||
if let Some(info) = route.get_peer_info(peer_id).await {
|
||||
return Some(info);
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub async fn get_origin_my_peer_id(
|
||||
&self,
|
||||
network_name: &str,
|
||||
foreign_my_peer_id: PeerId,
|
||||
) -> Option<PeerId> {
|
||||
for route in self.routes.read().await.iter() {
|
||||
let origin_peer_id = route
|
||||
.get_origin_my_peer_id(network_name, foreign_my_peer_id)
|
||||
.await;
|
||||
if origin_peer_id.is_some() {
|
||||
return origin_peer_id;
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.peer_map.is_empty()
|
||||
}
|
||||
|
||||
pub fn list_peers(&self) -> Vec<PeerId> {
|
||||
let mut ret = Vec::new();
|
||||
for item in self.peer_map.iter() {
|
||||
let peer_id = item.key();
|
||||
ret.push(*peer_id);
|
||||
}
|
||||
ret
|
||||
}
|
||||
|
||||
pub async fn list_peers_with_conn(&self) -> Vec<PeerId> {
|
||||
let mut ret = Vec::new();
|
||||
for item in self.peer_map.iter() {
|
||||
if item.value().has_live_conns() {
|
||||
ret.push(*item.key());
|
||||
}
|
||||
}
|
||||
ret
|
||||
}
|
||||
|
||||
pub async fn list_peer_conns(&self, peer_id: PeerId) -> Option<Vec<PeerConnInfo>> {
|
||||
if let Some(p) = self.get_peer_by_id(peer_id) {
|
||||
Some(p.list_peer_conns().await)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_peer_default_conn_id(&self, peer_id: PeerId) -> Option<PeerConnId> {
|
||||
self.get_peer_by_id(peer_id)
|
||||
.map(|p| p.get_default_conn_id())
|
||||
}
|
||||
|
||||
pub fn get_peer_identity_type(&self, peer_id: PeerId) -> Option<PeerIdentityType> {
|
||||
self.get_peer_by_id(peer_id)
|
||||
.and_then(|p| p.get_peer_identity_type())
|
||||
}
|
||||
|
||||
pub fn get_peer_public_key(&self, peer_id: PeerId) -> Option<Vec<u8>> {
|
||||
self.get_peer_by_id(peer_id)
|
||||
.and_then(|p| p.get_peer_public_key())
|
||||
}
|
||||
|
||||
pub async fn close_peer_conn(
|
||||
&self,
|
||||
peer_id: PeerId,
|
||||
conn_id: &PeerConnId,
|
||||
) -> Result<(), Error> {
|
||||
if let Some(p) = self.get_peer_by_id(peer_id) {
|
||||
p.close_peer_conn(conn_id).await
|
||||
} else {
|
||||
Err(Error::NotFound)
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn close_peer(&self, peer_id: PeerId) -> Result<(), TunnelError> {
|
||||
let remove_ret = self.peer_map.remove(&peer_id);
|
||||
shrink_dashmap(&self.peer_map, None);
|
||||
|
||||
self.context.issue_event(PeerEvent::PeerRemoved(peer_id));
|
||||
tracing::info!(
|
||||
?peer_id,
|
||||
has_old_value = ?remove_ret.is_some(),
|
||||
peer_ref_counter = ?remove_ret.map(|v| Arc::strong_count(&v.1)),
|
||||
"peer is closed"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn add_route(&self, route: ArcRoute) {
|
||||
let mut routes = self.routes.write().await;
|
||||
routes.insert(0, route);
|
||||
}
|
||||
|
||||
pub(crate) async fn clear_resources(&self) {
|
||||
for peer_id in self.list_peers() {
|
||||
let _ = self.close_peer(peer_id).await;
|
||||
}
|
||||
let routes = {
|
||||
let mut routes = self.routes.write().await;
|
||||
std::mem::take(&mut *routes)
|
||||
};
|
||||
for route in routes {
|
||||
route.close().await;
|
||||
}
|
||||
self.alive_client_urls.lock().clear();
|
||||
}
|
||||
|
||||
pub async fn clean_peer_without_conn(&self) {
|
||||
let mut to_remove = vec![];
|
||||
|
||||
for peer_id in self.list_peers() {
|
||||
let conns = self.list_peer_conns(peer_id).await;
|
||||
if conns.is_none() || conns.as_ref().unwrap().is_empty() {
|
||||
to_remove.push(peer_id);
|
||||
}
|
||||
}
|
||||
|
||||
for peer_id in to_remove {
|
||||
self.close_peer(peer_id).await.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn list_routes(&self) -> DashMap<PeerId, PeerId> {
|
||||
let route_map = DashMap::new();
|
||||
for route in self.routes.read().await.iter() {
|
||||
for item in route.list_routes().await.iter() {
|
||||
route_map.insert(item.peer_id, item.next_hop_peer_id);
|
||||
}
|
||||
}
|
||||
route_map
|
||||
}
|
||||
|
||||
pub async fn list_route_infos(&self) -> Vec<CoreRoute> {
|
||||
if let Some(route) = self.routes.read().await.iter().next() {
|
||||
return route.list_routes().await;
|
||||
}
|
||||
vec![]
|
||||
}
|
||||
|
||||
pub async fn need_relay_by_foreign_network(&self, dst_peer_id: PeerId) -> Result<bool, Error> {
|
||||
// if gateway_peer_id is not connected to me, means need relay by foreign network
|
||||
let gateway_id = self
|
||||
.get_gateway_peer_id(dst_peer_id, NextHopPolicy::LeastHop)
|
||||
.await
|
||||
.ok_or(Error::RouteError(Some(format!(
|
||||
"peer map need_relay_by_foreign_network no gateway for dst_peer_id: {}",
|
||||
dst_peer_id
|
||||
))))?;
|
||||
|
||||
Ok(!self.has_peer(gateway_id))
|
||||
}
|
||||
|
||||
pub fn my_peer_id(&self) -> PeerId {
|
||||
self.my_peer_id
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PeerMap {
|
||||
fn drop(&mut self) {
|
||||
tracing::debug!(
|
||||
self.my_peer_id,
|
||||
network = ?self.context.network_identity(),
|
||||
"PeerMap is dropped"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// Aggregates the directly-connected peers of several peer maps into the
|
||||
/// peer-center reporting format, keeping the lowest observed latency per peer.
|
||||
pub(crate) async fn direct_peer_info(peer_maps: &[Arc<PeerMap>]) -> PeerInfoForGlobalMap {
|
||||
let mut peers = BTreeSet::new();
|
||||
for peer_map in peer_maps {
|
||||
peers.extend(peer_map.list_peers());
|
||||
}
|
||||
|
||||
let mut ret = PeerInfoForGlobalMap::default();
|
||||
for peer in peers {
|
||||
let mut conns = None;
|
||||
for peer_map in peer_maps {
|
||||
if let Some(found) = peer_map.list_peer_conns(peer).await {
|
||||
conns = Some(found);
|
||||
break;
|
||||
}
|
||||
}
|
||||
let Some(min_lat) = conns
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.map(|conn| conn.stats.as_ref().unwrap().latency_us)
|
||||
.min()
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
ret.direct_peers.insert(
|
||||
peer,
|
||||
DirectConnectedPeerInfo {
|
||||
latency_ms: std::cmp::max(1, (min_lat as u32 / 1000) as i32),
|
||||
},
|
||||
);
|
||||
}
|
||||
ret
|
||||
}
|
||||
@@ -0,0 +1,563 @@
|
||||
use std::sync::{
|
||||
Arc, RwLock,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use anyhow::anyhow;
|
||||
use crossbeam::atomic::AtomicCell;
|
||||
use dashmap::DashMap;
|
||||
|
||||
use crate::peers::util::shrink_dashmap;
|
||||
use crate::tunnel::secure_datagram::{SecureDatagramDirection, SecureDatagramSession};
|
||||
use crate::{config::PeerId, packet::ZCPacket};
|
||||
|
||||
const SESSION_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
|
||||
|
||||
pub struct UpsertResponderSessionReturn {
|
||||
pub session: Arc<PeerSession>,
|
||||
pub action: PeerSessionAction,
|
||||
pub session_generation: u32,
|
||||
pub root_key: Option<[u8; 32]>,
|
||||
pub initial_epoch: u32,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum PeerSessionAction {
|
||||
Join,
|
||||
Sync,
|
||||
Create,
|
||||
}
|
||||
|
||||
#[derive(PartialEq, Clone, Eq, Hash, Debug)]
|
||||
pub struct SessionKey {
|
||||
network_name: String,
|
||||
peer_id: PeerId,
|
||||
}
|
||||
|
||||
impl SessionKey {
|
||||
pub fn new(network_name: String, peer_id: PeerId) -> Self {
|
||||
Self {
|
||||
network_name,
|
||||
peer_id,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PeerSessionStore {
|
||||
sessions: Arc<DashMap<SessionKey, PeerSessionEntry>>,
|
||||
}
|
||||
|
||||
struct PeerSessionEntry {
|
||||
session: Arc<PeerSession>,
|
||||
last_used_at: AtomicCell<Instant>,
|
||||
}
|
||||
|
||||
impl PeerSessionEntry {
|
||||
fn new(session: Arc<PeerSession>) -> Self {
|
||||
Self {
|
||||
session,
|
||||
last_used_at: AtomicCell::new(Instant::now()),
|
||||
}
|
||||
}
|
||||
|
||||
fn touch(&self) {
|
||||
self.last_used_at.store(Instant::now());
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for PeerSessionStore {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
sessions: Arc::new(DashMap::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerSessionStore {
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
|
||||
pub fn get(&self, key: &SessionKey) -> Option<Arc<PeerSession>> {
|
||||
let session = {
|
||||
let entry = self.sessions.get(key)?;
|
||||
entry.touch();
|
||||
entry.session.clone()
|
||||
};
|
||||
if session.is_valid() {
|
||||
Some(session)
|
||||
} else {
|
||||
self.sessions.remove(key);
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
pub fn remove(&self, key: &SessionKey) {
|
||||
self.sessions.remove(key);
|
||||
}
|
||||
|
||||
pub fn evict_unused_sessions(&self) {
|
||||
self.evict_unused_sessions_idle(SESSION_IDLE_TIMEOUT);
|
||||
}
|
||||
|
||||
pub fn evict_unused_sessions_idle(&self, idle: Duration) {
|
||||
let now = Instant::now();
|
||||
self.sessions.retain(|_key, entry| {
|
||||
entry.session.is_valid()
|
||||
&& (Arc::strong_count(&entry.session) > 1
|
||||
|| now.saturating_duration_since(entry.last_used_at.load()) < idle)
|
||||
});
|
||||
shrink_dashmap(&self.sessions, None);
|
||||
}
|
||||
|
||||
#[tracing::instrument(skip(self))]
|
||||
pub fn upsert_responder_session(
|
||||
&self,
|
||||
key: &SessionKey,
|
||||
a_session_generation: Option<u32>,
|
||||
send_algorithm: String,
|
||||
recv_algorithm: String,
|
||||
peer_static_pubkey: Option<[u8; 32]>,
|
||||
) -> Result<UpsertResponderSessionReturn, anyhow::Error> {
|
||||
tracing::event!(tracing::Level::INFO, ?key, "upsert_responder_session");
|
||||
let existing = self
|
||||
.sessions
|
||||
.get(key)
|
||||
.map(|v| {
|
||||
v.touch();
|
||||
v.session.clone()
|
||||
})
|
||||
.filter(|s| s.is_valid());
|
||||
match existing {
|
||||
None => {
|
||||
let root_key = PeerSession::new_root_key();
|
||||
let session_generation = 1u32;
|
||||
let initial_epoch = 0u32;
|
||||
let session = Arc::new(PeerSession::new(
|
||||
key.peer_id,
|
||||
root_key,
|
||||
session_generation,
|
||||
initial_epoch,
|
||||
send_algorithm,
|
||||
recv_algorithm,
|
||||
peer_static_pubkey,
|
||||
));
|
||||
self.sessions
|
||||
.insert(key.clone(), PeerSessionEntry::new(session.clone()));
|
||||
Ok(UpsertResponderSessionReturn {
|
||||
session,
|
||||
action: PeerSessionAction::Create,
|
||||
session_generation,
|
||||
root_key: Some(root_key),
|
||||
initial_epoch,
|
||||
})
|
||||
}
|
||||
Some(session) => {
|
||||
session.check_encrypt_algo_same(&send_algorithm, &recv_algorithm)?;
|
||||
session.check_or_set_peer_static_pubkey(peer_static_pubkey)?;
|
||||
let local_gen = session.session_generation();
|
||||
if a_session_generation.is_some_and(|g| g == local_gen) {
|
||||
Ok(UpsertResponderSessionReturn {
|
||||
session,
|
||||
action: PeerSessionAction::Join,
|
||||
session_generation: local_gen,
|
||||
root_key: None,
|
||||
initial_epoch: 0,
|
||||
})
|
||||
} else {
|
||||
let initial_epoch = session.next_sync_epoch();
|
||||
let root_key = session.root_key();
|
||||
Ok(UpsertResponderSessionReturn {
|
||||
session,
|
||||
action: PeerSessionAction::Sync,
|
||||
session_generation: local_gen,
|
||||
root_key: Some(root_key),
|
||||
initial_epoch,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
#[tracing::instrument(skip(self))]
|
||||
pub fn apply_initiator_action(
|
||||
&self,
|
||||
key: &SessionKey,
|
||||
action: PeerSessionAction,
|
||||
b_session_generation: u32,
|
||||
root_key_32: Option<[u8; 32]>,
|
||||
initial_epoch: u32,
|
||||
send_algorithm: String,
|
||||
recv_algorithm: String,
|
||||
peer_static_pubkey: Option<[u8; 32]>,
|
||||
) -> Result<Arc<PeerSession>, anyhow::Error> {
|
||||
tracing::event!(tracing::Level::INFO, "apply_initiator_action {:?}", key);
|
||||
match action {
|
||||
PeerSessionAction::Join => {
|
||||
let Some(session) = self.get(key) else {
|
||||
return Err(anyhow!("no local session for JOIN"));
|
||||
};
|
||||
session.check_encrypt_algo_same(&send_algorithm, &recv_algorithm)?;
|
||||
session.check_or_set_peer_static_pubkey(peer_static_pubkey)?;
|
||||
if session.session_generation() != b_session_generation {
|
||||
return Err(anyhow!("JOIN generation mismatch"));
|
||||
}
|
||||
Ok(session)
|
||||
}
|
||||
PeerSessionAction::Sync | PeerSessionAction::Create => {
|
||||
let root_key = root_key_32.ok_or_else(|| anyhow!("missing root_key"))?;
|
||||
if let Some(existing) = self.sessions.get(key)
|
||||
&& !existing.session.is_valid()
|
||||
{
|
||||
drop(existing);
|
||||
self.sessions.remove(key);
|
||||
}
|
||||
let session = {
|
||||
let entry = self.sessions.entry(key.clone()).or_insert_with(|| {
|
||||
PeerSessionEntry::new(Arc::new(PeerSession::new(
|
||||
key.peer_id,
|
||||
root_key,
|
||||
b_session_generation,
|
||||
initial_epoch,
|
||||
send_algorithm.clone(),
|
||||
recv_algorithm.clone(),
|
||||
peer_static_pubkey,
|
||||
)))
|
||||
});
|
||||
entry.touch();
|
||||
entry.session.clone()
|
||||
};
|
||||
session.check_encrypt_algo_same(&send_algorithm, &recv_algorithm)?;
|
||||
session.check_or_set_peer_static_pubkey(peer_static_pubkey)?;
|
||||
session.sync_root_key(
|
||||
root_key,
|
||||
b_session_generation,
|
||||
initial_epoch,
|
||||
matches!(action, PeerSessionAction::Sync),
|
||||
);
|
||||
Ok(session)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PeerSession {
|
||||
peer_id: PeerId,
|
||||
peer_static_pubkey: RwLock<Option<[u8; 32]>>,
|
||||
datagram: SecureDatagramSession,
|
||||
invalidated: AtomicBool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PeerSession {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("PeerSession")
|
||||
.field("peer_id", &self.peer_id)
|
||||
.field("peer_static_pubkey", &self.peer_static_pubkey)
|
||||
.field("datagram", &self.datagram)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerSession {
|
||||
pub fn new(
|
||||
peer_id: PeerId,
|
||||
root_key: [u8; 32],
|
||||
session_generation: u32,
|
||||
initial_epoch: u32,
|
||||
send_cipher_algorithm: String,
|
||||
recv_cipher_algorithm: String,
|
||||
peer_static_pubkey: Option<[u8; 32]>,
|
||||
) -> Self {
|
||||
Self {
|
||||
peer_id,
|
||||
peer_static_pubkey: RwLock::new(peer_static_pubkey),
|
||||
datagram: SecureDatagramSession::new(
|
||||
root_key,
|
||||
session_generation,
|
||||
initial_epoch,
|
||||
send_cipher_algorithm,
|
||||
recv_cipher_algorithm,
|
||||
),
|
||||
invalidated: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn peer_id(&self) -> PeerId {
|
||||
self.peer_id
|
||||
}
|
||||
|
||||
pub fn invalidate(&self) {
|
||||
self.invalidated.store(true, Ordering::Relaxed);
|
||||
self.datagram.invalidate();
|
||||
}
|
||||
|
||||
pub fn is_valid(&self) -> bool {
|
||||
!self.invalidated.load(Ordering::Relaxed) && self.datagram.is_valid()
|
||||
}
|
||||
|
||||
pub fn session_generation(&self) -> u32 {
|
||||
self.datagram.session_generation()
|
||||
}
|
||||
|
||||
pub fn root_key(&self) -> [u8; 32] {
|
||||
self.datagram.root_key()
|
||||
}
|
||||
|
||||
pub fn new_root_key() -> [u8; 32] {
|
||||
SecureDatagramSession::new_root_key()
|
||||
}
|
||||
|
||||
pub fn next_sync_epoch(&self) -> u32 {
|
||||
self.datagram.next_sync_epoch()
|
||||
}
|
||||
|
||||
pub fn check_encrypt_algo_same(
|
||||
&self,
|
||||
send_algorithm: &str,
|
||||
recv_algorithm: &str,
|
||||
) -> Result<(), anyhow::Error> {
|
||||
self.datagram
|
||||
.check_encrypt_algo_same(send_algorithm, recv_algorithm)
|
||||
}
|
||||
|
||||
pub fn check_or_set_peer_static_pubkey(
|
||||
&self,
|
||||
peer_static_pubkey: Option<[u8; 32]>,
|
||||
) -> Result<(), anyhow::Error> {
|
||||
let Some(peer_static_pubkey) = peer_static_pubkey else {
|
||||
return Ok(());
|
||||
};
|
||||
let mut guard = self.peer_static_pubkey.write().unwrap();
|
||||
if let Some(existing) = *guard {
|
||||
if existing != peer_static_pubkey {
|
||||
return Err(anyhow!("peer static pubkey mismatch"));
|
||||
}
|
||||
return Ok(());
|
||||
}
|
||||
*guard = Some(peer_static_pubkey);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn sync_root_key(
|
||||
&self,
|
||||
root_key: [u8; 32],
|
||||
session_generation: u32,
|
||||
initial_epoch: u32,
|
||||
preserve_rx_grace: bool,
|
||||
) {
|
||||
self.datagram.sync_root_key(
|
||||
root_key,
|
||||
session_generation,
|
||||
initial_epoch,
|
||||
preserve_rx_grace,
|
||||
);
|
||||
}
|
||||
|
||||
pub fn dir_for_sender(
|
||||
sender_peer_id: PeerId,
|
||||
receiver_peer_id: PeerId,
|
||||
) -> SecureDatagramDirection {
|
||||
if sender_peer_id < receiver_peer_id {
|
||||
SecureDatagramDirection::AToB
|
||||
} else {
|
||||
SecureDatagramDirection::BToA
|
||||
}
|
||||
}
|
||||
|
||||
pub fn encrypt_payload(
|
||||
&self,
|
||||
sender_peer_id: PeerId,
|
||||
receiver_peer_id: PeerId,
|
||||
pkt: &mut ZCPacket,
|
||||
) -> Result<(), anyhow::Error> {
|
||||
if !self.is_valid() {
|
||||
return Err(anyhow!("session invalidated"));
|
||||
}
|
||||
self.datagram
|
||||
.encrypt_payload(Self::dir_for_sender(sender_peer_id, receiver_peer_id), pkt)
|
||||
}
|
||||
|
||||
pub fn decrypt_payload(
|
||||
&self,
|
||||
sender_peer_id: PeerId,
|
||||
receiver_peer_id: PeerId,
|
||||
ciphertext_with_tail: &mut ZCPacket,
|
||||
) -> Result<(), anyhow::Error> {
|
||||
if !self.is_valid() {
|
||||
return Err(anyhow!("session invalidated"));
|
||||
}
|
||||
self.datagram.decrypt_payload(
|
||||
Self::dir_for_sender(sender_peer_id, receiver_peer_id),
|
||||
ciphertext_with_tail,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-utils"))]
|
||||
mod test_utils {
|
||||
use super::*;
|
||||
|
||||
impl PeerSessionStore {
|
||||
#[doc(hidden)]
|
||||
pub(crate) fn contains_valid(&self, key: &SessionKey) -> bool {
|
||||
self.sessions
|
||||
.get(key)
|
||||
.is_some_and(|entry| entry.session.is_valid())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
impl PeerSessionStore {
|
||||
fn insert_session(&self, key: SessionKey, session: Arc<PeerSession>) {
|
||||
self.sessions.insert(key, PeerSessionEntry::new(session));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[cfg(all(feature = "aes-gcm", feature = "chacha20"))]
|
||||
fn peer_session_supports_asymmetric_algorithms() {
|
||||
let a: PeerId = 10;
|
||||
let b: PeerId = 20;
|
||||
let root_key = PeerSession::new_root_key();
|
||||
let generation = 1u32;
|
||||
let initial_epoch = 0u32;
|
||||
|
||||
let sa = PeerSession::new(
|
||||
b,
|
||||
root_key,
|
||||
generation,
|
||||
initial_epoch,
|
||||
"aes-256-gcm".to_string(),
|
||||
"chacha20-poly1305".to_string(),
|
||||
None,
|
||||
);
|
||||
let sb = PeerSession::new(
|
||||
a,
|
||||
root_key,
|
||||
generation,
|
||||
initial_epoch,
|
||||
"chacha20-poly1305".to_string(),
|
||||
"aes-256-gcm".to_string(),
|
||||
None,
|
||||
);
|
||||
|
||||
let plaintext1 = b"hello from a";
|
||||
let mut pkt1 = ZCPacket::new_with_payload(plaintext1);
|
||||
pkt1.fill_peer_manager_hdr(a as u32, b as u32, 0);
|
||||
sa.encrypt_payload(a, b, &mut pkt1).unwrap();
|
||||
sb.decrypt_payload(a, b, &mut pkt1).unwrap();
|
||||
assert_eq!(pkt1.payload(), plaintext1);
|
||||
|
||||
let plaintext2 = b"hello from b";
|
||||
let mut pkt2 = ZCPacket::new_with_payload(plaintext2);
|
||||
pkt2.fill_peer_manager_hdr(b as u32, a as u32, 0);
|
||||
sb.encrypt_payload(b, a, &mut pkt2).unwrap();
|
||||
sa.decrypt_payload(b, a, &mut pkt2).unwrap();
|
||||
assert_eq!(pkt2.payload(), plaintext2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn peer_session_store_keeps_recent_session_without_external_refs() {
|
||||
let store = PeerSessionStore::new();
|
||||
let key = SessionKey::new("net".to_string(), 20);
|
||||
let session = Arc::new(PeerSession::new(
|
||||
20,
|
||||
PeerSession::new_root_key(),
|
||||
1,
|
||||
0,
|
||||
"aes-gcm".to_string(),
|
||||
"aes-gcm".to_string(),
|
||||
None,
|
||||
));
|
||||
store.insert_session(key.clone(), session);
|
||||
|
||||
assert!(store.get(&key).is_some());
|
||||
store.evict_unused_sessions();
|
||||
|
||||
assert!(
|
||||
store.get(&key).is_some(),
|
||||
"recent relay sessions should survive the periodic GC"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn peer_session_store_evicts_idle_session_without_external_refs() {
|
||||
let store = PeerSessionStore::new();
|
||||
let key = SessionKey::new("net".to_string(), 20);
|
||||
let session = Arc::new(PeerSession::new(
|
||||
20,
|
||||
PeerSession::new_root_key(),
|
||||
1,
|
||||
0,
|
||||
"aes-gcm".to_string(),
|
||||
"aes-gcm".to_string(),
|
||||
None,
|
||||
));
|
||||
store.insert_session(key.clone(), session);
|
||||
|
||||
store.evict_unused_sessions_idle(Duration::from_millis(0));
|
||||
|
||||
assert!(
|
||||
store.get(&key).is_none(),
|
||||
"idle sessions without external users should still be collected"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn peer_session_store_evicts_invalid_recent_session() {
|
||||
let store = PeerSessionStore::new();
|
||||
let key = SessionKey::new("net".to_string(), 20);
|
||||
let session = Arc::new(PeerSession::new(
|
||||
20,
|
||||
PeerSession::new_root_key(),
|
||||
1,
|
||||
0,
|
||||
"aes-gcm".to_string(),
|
||||
"aes-gcm".to_string(),
|
||||
None,
|
||||
));
|
||||
store.insert_session(key.clone(), session);
|
||||
|
||||
let session = store.get(&key).unwrap();
|
||||
session.invalidate();
|
||||
drop(session);
|
||||
store.evict_unused_sessions();
|
||||
|
||||
assert!(
|
||||
!store.sessions.contains_key(&key),
|
||||
"invalid sessions should not be kept by recent activity"
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(feature = "test-utils")]
|
||||
#[test]
|
||||
fn contains_valid_does_not_refresh_session_activity() {
|
||||
let store = PeerSessionStore::new();
|
||||
let key = SessionKey::new("net".to_string(), 20);
|
||||
let session = Arc::new(PeerSession::new(
|
||||
20,
|
||||
PeerSession::new_root_key(),
|
||||
1,
|
||||
0,
|
||||
"aes-gcm".to_string(),
|
||||
"aes-gcm".to_string(),
|
||||
None,
|
||||
));
|
||||
store.insert_session(key.clone(), session);
|
||||
let last_used_at = store.sessions.get(&key).unwrap().last_used_at.load();
|
||||
|
||||
assert!(store.contains_valid(&key));
|
||||
|
||||
assert_eq!(
|
||||
store.sessions.get(&key).unwrap().last_used_at.load(),
|
||||
last_used_at
|
||||
);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,497 @@
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
sync::{Arc, Mutex},
|
||||
time::{Duration, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use base64::{Engine, engine::general_purpose::STANDARD as BASE64_STANDARD};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use x25519_dalek::{PublicKey, StaticSecret};
|
||||
|
||||
use crate::proto::peer_rpc::{TrustedCredentialPubkey, TrustedCredentialPubkeyProof};
|
||||
|
||||
fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn current_unix_timestamp() -> i64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs() as i64
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CredentialCreateOptions {
|
||||
pub groups: Vec<String>,
|
||||
pub allow_relay: bool,
|
||||
pub allowed_proxy_cidrs: Vec<String>,
|
||||
pub ttl: Duration,
|
||||
pub credential_id: Option<String>,
|
||||
pub reusable: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub(crate) struct CredentialEntry {
|
||||
pubkey: String,
|
||||
#[serde(default)]
|
||||
secret: String,
|
||||
groups: Vec<String>,
|
||||
allow_relay: bool,
|
||||
allowed_proxy_cidrs: Vec<String>,
|
||||
#[serde(default = "default_true")]
|
||||
reusable: bool,
|
||||
expiry_unix: i64,
|
||||
created_at_unix: i64,
|
||||
}
|
||||
|
||||
impl CredentialEntry {
|
||||
fn is_active_at(&self, now: i64) -> bool {
|
||||
self.expiry_unix > now
|
||||
}
|
||||
|
||||
fn to_trusted_credential(&self) -> Option<TrustedCredentialPubkey> {
|
||||
Some(TrustedCredentialPubkey {
|
||||
pubkey: CredentialManager::decode_pubkey_b64(&self.pubkey)?,
|
||||
groups: self.groups.clone(),
|
||||
allow_relay: self.allow_relay,
|
||||
expiry_unix: self.expiry_unix,
|
||||
allowed_proxy_cidrs: self.allowed_proxy_cidrs.clone(),
|
||||
reusable: Some(self.reusable),
|
||||
})
|
||||
}
|
||||
|
||||
fn to_credential_info(&self, credential_id: &str) -> CredentialInfo {
|
||||
CredentialInfo {
|
||||
credential_id: credential_id.to_string(),
|
||||
groups: self.groups.clone(),
|
||||
allow_relay: self.allow_relay,
|
||||
expiry_unix: self.expiry_unix,
|
||||
allowed_proxy_cidrs: self.allowed_proxy_cidrs.clone(),
|
||||
reusable: Some(self.reusable),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct CredentialInfo {
|
||||
pub credential_id: String,
|
||||
pub groups: Vec<String>,
|
||||
pub allow_relay: bool,
|
||||
pub expiry_unix: i64,
|
||||
pub allowed_proxy_cidrs: Vec<String>,
|
||||
pub reusable: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct GeneratedCredential {
|
||||
pub credential_id: String,
|
||||
pub secret: String,
|
||||
pub changed: bool,
|
||||
}
|
||||
|
||||
pub trait CredentialStorage: Send + Sync + 'static {
|
||||
fn load(&self) -> anyhow::Result<Option<String>>;
|
||||
fn store(&self, serialized_credentials: &str) -> anyhow::Result<()>;
|
||||
}
|
||||
|
||||
pub(crate) struct CredentialManager {
|
||||
credentials: Mutex<HashMap<String, CredentialEntry>>,
|
||||
storage: Option<Arc<dyn CredentialStorage>>,
|
||||
storage_write: Mutex<()>,
|
||||
}
|
||||
|
||||
impl Default for CredentialManager {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl CredentialManager {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
credentials: Mutex::new(HashMap::new()),
|
||||
storage: None,
|
||||
storage_write: Mutex::new(()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_storage(storage: Arc<dyn CredentialStorage>) -> Self {
|
||||
let credentials = match storage.load() {
|
||||
Ok(Some(serialized)) => serde_json::from_str(&serialized).unwrap_or_else(|error| {
|
||||
tracing::warn!(?error, "failed to parse stored credentials");
|
||||
HashMap::new()
|
||||
}),
|
||||
Ok(None) => HashMap::new(),
|
||||
Err(error) => {
|
||||
tracing::warn!(?error, "failed to load stored credentials");
|
||||
HashMap::new()
|
||||
}
|
||||
};
|
||||
Self {
|
||||
credentials: Mutex::new(credentials),
|
||||
storage: Some(storage),
|
||||
storage_write: Mutex::new(()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_entries<R>(&self, f: impl FnOnce(&HashMap<String, CredentialEntry>) -> R) -> R {
|
||||
let credentials = self.credentials.lock().unwrap();
|
||||
f(&credentials)
|
||||
}
|
||||
|
||||
pub fn generate_credential_with_options(
|
||||
&self,
|
||||
groups: Vec<String>,
|
||||
allow_relay: bool,
|
||||
allowed_proxy_cidrs: Vec<String>,
|
||||
ttl: Duration,
|
||||
credential_id: Option<String>,
|
||||
reusable: bool,
|
||||
) -> GeneratedCredential {
|
||||
self.remove_expired_credentials();
|
||||
self.generate_credential_with_options_after_cleanup(
|
||||
groups,
|
||||
allow_relay,
|
||||
allowed_proxy_cidrs,
|
||||
ttl,
|
||||
credential_id,
|
||||
reusable,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn generate_credential_with_options_after_cleanup(
|
||||
&self,
|
||||
groups: Vec<String>,
|
||||
allow_relay: bool,
|
||||
allowed_proxy_cidrs: Vec<String>,
|
||||
ttl: Duration,
|
||||
credential_id: Option<String>,
|
||||
reusable: bool,
|
||||
) -> GeneratedCredential {
|
||||
let generated = {
|
||||
let mut credentials = self.credentials.lock().unwrap();
|
||||
let id = if let Some(id) = credential_id
|
||||
.map(|x| x.trim().to_string())
|
||||
.filter(|x| !x.is_empty())
|
||||
{
|
||||
if let Some(existing) = credentials.get(&id)
|
||||
&& !existing.secret.is_empty()
|
||||
{
|
||||
return GeneratedCredential {
|
||||
credential_id: id,
|
||||
secret: existing.secret.clone(),
|
||||
changed: false,
|
||||
};
|
||||
}
|
||||
id
|
||||
} else {
|
||||
uuid::Uuid::new_v4().to_string()
|
||||
};
|
||||
|
||||
let (entry, secret) =
|
||||
Self::build_entry(groups, allow_relay, allowed_proxy_cidrs, reusable, ttl);
|
||||
credentials.insert(id.clone(), entry);
|
||||
GeneratedCredential {
|
||||
credential_id: id,
|
||||
secret,
|
||||
changed: true,
|
||||
}
|
||||
};
|
||||
self.persist();
|
||||
generated
|
||||
}
|
||||
|
||||
fn build_entry(
|
||||
groups: Vec<String>,
|
||||
allow_relay: bool,
|
||||
allowed_proxy_cidrs: Vec<String>,
|
||||
reusable: bool,
|
||||
ttl: Duration,
|
||||
) -> (CredentialEntry, String) {
|
||||
let private = StaticSecret::random_from_rng(rand::rngs::OsRng);
|
||||
let public = PublicKey::from(&private);
|
||||
let pubkey = BASE64_STANDARD.encode(public.as_bytes());
|
||||
let secret = BASE64_STANDARD.encode(private.as_bytes());
|
||||
|
||||
let now = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap()
|
||||
.as_secs() as i64;
|
||||
let expiry_unix = now + ttl.as_secs() as i64;
|
||||
|
||||
let entry = CredentialEntry {
|
||||
pubkey,
|
||||
secret: secret.clone(),
|
||||
groups,
|
||||
allow_relay,
|
||||
allowed_proxy_cidrs,
|
||||
reusable,
|
||||
expiry_unix,
|
||||
created_at_unix: now,
|
||||
};
|
||||
(entry, secret)
|
||||
}
|
||||
|
||||
pub fn revoke_credential(&self, credential_id: &str) -> bool {
|
||||
let removed = self
|
||||
.credentials
|
||||
.lock()
|
||||
.unwrap()
|
||||
.remove(credential_id)
|
||||
.is_some();
|
||||
if removed {
|
||||
self.persist();
|
||||
}
|
||||
removed
|
||||
}
|
||||
|
||||
pub fn remove_expired_credentials(&self) -> bool {
|
||||
self.remove_expired_credentials_at(current_unix_timestamp())
|
||||
}
|
||||
|
||||
fn remove_expired_credentials_at(&self, now: i64) -> bool {
|
||||
let mut credentials = self.credentials.lock().unwrap();
|
||||
let before = credentials.len();
|
||||
credentials.retain(|_, entry| entry.is_active_at(now));
|
||||
let changed = before != credentials.len();
|
||||
drop(credentials);
|
||||
if changed {
|
||||
self.persist();
|
||||
}
|
||||
changed
|
||||
}
|
||||
|
||||
pub fn get_trusted_pubkeys(&self, network_secret: &str) -> Vec<TrustedCredentialPubkeyProof> {
|
||||
let now = current_unix_timestamp();
|
||||
|
||||
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()
|
||||
}
|
||||
|
||||
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))
|
||||
}
|
||||
|
||||
pub fn list_credentials(&self) -> Vec<CredentialInfo> {
|
||||
let now = current_unix_timestamp();
|
||||
|
||||
self.credentials
|
||||
.lock()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.filter(|(_, entry)| entry.is_active_at(now))
|
||||
.map(|(id, entry)| entry.to_credential_info(id))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn decode_pubkey_b64(s: &str) -> Option<Vec<u8>> {
|
||||
let decoded = BASE64_STANDARD.decode(s).ok()?;
|
||||
if decoded.len() != 32 {
|
||||
return None;
|
||||
}
|
||||
Some(decoded)
|
||||
}
|
||||
|
||||
fn persist(&self) {
|
||||
let Some(storage) = &self.storage else {
|
||||
return;
|
||||
};
|
||||
let _storage_write = self.storage_write.lock().unwrap();
|
||||
let serialized = match self.with_entries(serde_json::to_string_pretty) {
|
||||
Ok(serialized) => serialized,
|
||||
Err(error) => {
|
||||
tracing::warn!(?error, "failed to serialize credentials");
|
||||
return;
|
||||
}
|
||||
};
|
||||
if let Err(error) = storage.store(&serialized) {
|
||||
tracing::warn!(?error, "failed to store credentials");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
impl CredentialManager {
|
||||
pub(crate) fn generate_credential(
|
||||
&self,
|
||||
groups: Vec<String>,
|
||||
allow_relay: bool,
|
||||
allowed_proxy_cidrs: Vec<String>,
|
||||
ttl: Duration,
|
||||
) -> GeneratedCredential {
|
||||
self.generate_credential_with_options(
|
||||
groups,
|
||||
allow_relay,
|
||||
allowed_proxy_cidrs,
|
||||
ttl,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
}
|
||||
|
||||
fn generate_credential_with_id(
|
||||
&self,
|
||||
groups: Vec<String>,
|
||||
allow_relay: bool,
|
||||
allowed_proxy_cidrs: Vec<String>,
|
||||
ttl: Duration,
|
||||
credential_id: Option<String>,
|
||||
) -> GeneratedCredential {
|
||||
self.generate_credential_with_options(
|
||||
groups,
|
||||
allow_relay,
|
||||
allowed_proxy_cidrs,
|
||||
ttl,
|
||||
credential_id,
|
||||
true,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct MemoryCredentialStorage {
|
||||
serialized: Mutex<Option<String>>,
|
||||
}
|
||||
|
||||
impl CredentialStorage for MemoryCredentialStorage {
|
||||
fn load(&self) -> anyhow::Result<Option<String>> {
|
||||
Ok(self.serialized.lock().unwrap().clone())
|
||||
}
|
||||
|
||||
fn store(&self, serialized_credentials: &str) -> anyhow::Result<()> {
|
||||
*self.serialized.lock().unwrap() = Some(serialized_credentials.to_owned());
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generate_and_revoke_credential() {
|
||||
let mgr = CredentialManager::new();
|
||||
let generated = mgr.generate_credential(
|
||||
vec!["guest".to_string()],
|
||||
false,
|
||||
vec![],
|
||||
Duration::from_secs(3600),
|
||||
);
|
||||
|
||||
assert!(!generated.credential_id.is_empty());
|
||||
assert!(!generated.secret.is_empty());
|
||||
assert!(generated.changed);
|
||||
assert!(uuid::Uuid::parse_str(&generated.credential_id).is_ok());
|
||||
|
||||
let privkey_bytes: [u8; 32] = BASE64_STANDARD
|
||||
.decode(&generated.secret)
|
||||
.unwrap()
|
||||
.try_into()
|
||||
.unwrap();
|
||||
let private = StaticSecret::from(privkey_bytes);
|
||||
let pubkey_bytes = PublicKey::from(&private).as_bytes().to_vec();
|
||||
assert!(mgr.is_pubkey_trusted(&pubkey_bytes));
|
||||
|
||||
let trusted = mgr.get_trusted_pubkeys("sec");
|
||||
assert_eq!(trusted.len(), 1);
|
||||
assert_eq!(
|
||||
trusted[0].credential.as_ref().unwrap().groups,
|
||||
vec!["guest".to_string()]
|
||||
);
|
||||
assert_eq!(trusted[0].credential.as_ref().unwrap().reusable, Some(true));
|
||||
|
||||
assert!(mgr.revoke_credential(&generated.credential_id));
|
||||
assert!(!mgr.is_pubkey_trusted(&pubkey_bytes));
|
||||
assert!(mgr.get_trusted_pubkeys("sec").is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn fixed_id_reuses_existing_secret() {
|
||||
let mgr = CredentialManager::new();
|
||||
let fixed_id = "fixed-credential-id".to_string();
|
||||
let first = mgr.generate_credential_with_id(
|
||||
vec!["group-a".to_string()],
|
||||
false,
|
||||
vec!["10.0.0.0/24".to_string()],
|
||||
Duration::from_secs(3600),
|
||||
Some(fixed_id.clone()),
|
||||
);
|
||||
let second = mgr.generate_credential_with_id(
|
||||
vec!["group-b".to_string()],
|
||||
true,
|
||||
vec!["192.168.0.0/16".to_string()],
|
||||
Duration::from_secs(7200),
|
||||
Some(fixed_id.clone()),
|
||||
);
|
||||
|
||||
assert_eq!(first.credential_id, fixed_id);
|
||||
assert_eq!(second.credential_id, fixed_id);
|
||||
assert_eq!(first.secret, second.secret);
|
||||
assert!(first.changed);
|
||||
assert!(!second.changed);
|
||||
|
||||
let list = mgr.list_credentials();
|
||||
assert_eq!(list.len(), 1);
|
||||
assert_eq!(list[0].credential_id, fixed_id);
|
||||
assert_eq!(list[0].groups, vec!["group-a".to_string()]);
|
||||
assert!(!list[0].allow_relay);
|
||||
assert_eq!(list[0].allowed_proxy_cidrs, vec!["10.0.0.0/24".to_string()]);
|
||||
assert_eq!(list[0].reusable, Some(true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn expired_credentials_are_filtered() {
|
||||
let mgr = CredentialManager::new();
|
||||
mgr.generate_credential(vec![], false, vec![], Duration::from_secs(3600));
|
||||
mgr.generate_credential(vec![], false, vec![], Duration::from_secs(0));
|
||||
|
||||
assert_eq!(mgr.list_credentials().len(), 1);
|
||||
assert!(mgr.remove_expired_credentials());
|
||||
assert_eq!(mgr.list_credentials().len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn injected_storage_loads_and_persists_mutations() {
|
||||
let storage = Arc::new(MemoryCredentialStorage::default());
|
||||
let manager = CredentialManager::from_storage(storage.clone());
|
||||
let generated =
|
||||
manager.generate_credential(vec![], false, vec![], Duration::from_secs(3600));
|
||||
|
||||
let reloaded = CredentialManager::from_storage(storage.clone());
|
||||
assert_eq!(
|
||||
reloaded.list_credentials()[0].credential_id,
|
||||
generated.credential_id
|
||||
);
|
||||
|
||||
assert!(manager.revoke_credential(&generated.credential_id));
|
||||
let reloaded = CredentialManager::from_storage(storage);
|
||||
assert!(reloaded.list_credentials().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn malformed_storage_starts_with_empty_credentials() {
|
||||
let storage = Arc::new(MemoryCredentialStorage {
|
||||
serialized: Mutex::new(Some("not json".to_owned())),
|
||||
});
|
||||
|
||||
let manager = CredentialManager::from_storage(storage);
|
||||
|
||||
assert!(manager.list_credentials().is_empty());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
use crate::{config::PeerId, tunnel::TunnelError};
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("wait response error: {0}")]
|
||||
WaitRespError(String),
|
||||
#[error("secret key error: {0}")]
|
||||
SecretKeyError(String),
|
||||
#[error("peer has no connection: {0}")]
|
||||
PeerNoConnectionError(PeerId),
|
||||
#[error("route error: {0:?}")]
|
||||
RouteError(Option<String>),
|
||||
#[error("not found")]
|
||||
NotFound,
|
||||
#[error(transparent)]
|
||||
Tunnel(#[from] TunnelError),
|
||||
#[error(transparent)]
|
||||
Other(#[from] anyhow::Error),
|
||||
}
|
||||
|
||||
impl From<snow::Error> for Error {
|
||||
fn from(value: snow::Error) -> Self {
|
||||
Self::WaitRespError(value.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl From<crate::foundation::time::error::Elapsed> for Error {
|
||||
fn from(value: crate::foundation::time::error::Elapsed) -> Self {
|
||||
Self::WaitRespError(value.to_string())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use tokio_util::task::AbortOnDropHandle;
|
||||
|
||||
use crate::{config::PeerId, packet::ZCPacket};
|
||||
|
||||
use crate::peers::{
|
||||
PacketRecvChan,
|
||||
conn::{peer_conn::PeerConn, peer_map::PeerMap},
|
||||
context::ArcPeerContext,
|
||||
error::Error,
|
||||
};
|
||||
|
||||
pub struct ForeignNetworkClient {
|
||||
peer_map: Arc<PeerMap>,
|
||||
task: Mutex<Option<AbortOnDropHandle<()>>>,
|
||||
}
|
||||
|
||||
impl ForeignNetworkClient {
|
||||
pub(crate) fn new(
|
||||
context: ArcPeerContext,
|
||||
packet_sender_to_mgr: PacketRecvChan,
|
||||
my_peer_id: PeerId,
|
||||
) -> Self {
|
||||
let peer_map = Arc::new(PeerMap::new(
|
||||
packet_sender_to_mgr,
|
||||
context.clone(),
|
||||
my_peer_id,
|
||||
));
|
||||
Self {
|
||||
peer_map,
|
||||
task: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn add_new_peer_conn(&self, peer_conn: PeerConn) -> Result<(), Error> {
|
||||
tracing::warn!(peer_conn = ?peer_conn.get_conn_info(), network = ?peer_conn.get_network_identity(), "add new peer conn in foreign network client");
|
||||
self.peer_map.add_new_peer_conn(peer_conn).await
|
||||
}
|
||||
|
||||
pub fn is_client_url_alive(&self, url: &url::Url) -> bool {
|
||||
self.peer_map.is_client_url_alive(url)
|
||||
}
|
||||
|
||||
pub fn has_next_hop(&self, peer_id: PeerId) -> bool {
|
||||
self.get_next_hop(peer_id).is_some()
|
||||
}
|
||||
|
||||
pub async fn list_public_peers(&self) -> Vec<PeerId> {
|
||||
self.peer_map.list_peers()
|
||||
}
|
||||
|
||||
pub fn get_next_hop(&self, peer_id: PeerId) -> Option<PeerId> {
|
||||
if self.peer_map.has_peer(peer_id) {
|
||||
return Some(peer_id);
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub async fn send_msg(&self, msg: ZCPacket, peer_id: PeerId) -> Result<(), Error> {
|
||||
if let Some(next_hop) = self.get_next_hop(peer_id) {
|
||||
let ret = self.peer_map.send_msg_directly(msg, next_hop).await;
|
||||
if ret.is_err() {
|
||||
tracing::error!(
|
||||
?ret,
|
||||
?peer_id,
|
||||
?next_hop,
|
||||
"foreign network client send msg failed"
|
||||
);
|
||||
} else {
|
||||
tracing::info!(
|
||||
?peer_id,
|
||||
?next_hop,
|
||||
"foreign network client send msg success"
|
||||
);
|
||||
}
|
||||
return ret;
|
||||
}
|
||||
Err(Error::RouteError(Some("no next hop".to_string())))
|
||||
}
|
||||
|
||||
pub async fn run(&self) {
|
||||
let peer_map = Arc::downgrade(&self.peer_map);
|
||||
*self.task.lock().unwrap() = Some(AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
loop {
|
||||
crate::foundation::time::sleep(crate::foundation::time::Duration::from_secs(1))
|
||||
.await;
|
||||
let Some(peer_map) = peer_map.upgrade() else {
|
||||
break;
|
||||
};
|
||||
peer_map.clean_peer_without_conn().await;
|
||||
}
|
||||
})));
|
||||
}
|
||||
|
||||
pub fn get_peer_map(&self) -> Arc<PeerMap> {
|
||||
self.peer_map.clone()
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,60 @@
|
||||
pub(crate) mod acl;
|
||||
pub(crate) mod admission;
|
||||
pub(crate) mod conn;
|
||||
pub mod context;
|
||||
pub mod credential_manager;
|
||||
pub mod error;
|
||||
pub mod foreign_network;
|
||||
pub mod peer_center;
|
||||
pub mod peer_manager;
|
||||
pub(crate) mod peer_rpc;
|
||||
pub mod public_ipv6;
|
||||
pub(crate) mod relay_peer_map;
|
||||
pub(crate) mod route;
|
||||
pub(crate) mod traffic_metrics;
|
||||
mod util;
|
||||
pub(crate) mod whitelist;
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) mod test_support;
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
use crate::packet::ZCPacket;
|
||||
|
||||
pub type PacketRecvChan = tokio::sync::mpsc::Sender<ZCPacket>;
|
||||
pub type PacketRecvChanReceiver = tokio::sync::mpsc::Receiver<ZCPacket>;
|
||||
|
||||
pub fn create_packet_recv_chan() -> (PacketRecvChan, PacketRecvChanReceiver) {
|
||||
tokio::sync::mpsc::channel(128)
|
||||
}
|
||||
|
||||
pub async fn recv_packet_from_chan(
|
||||
packet_recv_chan_receiver: &mut PacketRecvChanReceiver,
|
||||
) -> Result<ZCPacket, anyhow::Error> {
|
||||
packet_recv_chan_receiver
|
||||
.recv()
|
||||
.await
|
||||
.ok_or(anyhow::anyhow!("recv_packet_from_chan failed"))
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(Arc)]
|
||||
pub trait PeerPacketFilter {
|
||||
async fn try_process_packet_from_peer(&self, zc_packet: ZCPacket) -> Option<ZCPacket> {
|
||||
Some(zc_packet)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(Arc)]
|
||||
pub trait NicPacketFilter {
|
||||
async fn try_process_packet_from_nic(&self, data: &mut ZCPacket) -> bool;
|
||||
|
||||
fn id(&self) -> String {
|
||||
format!("{:p}", self)
|
||||
}
|
||||
}
|
||||
|
||||
pub type BoxPeerPacketFilter = Box<dyn PeerPacketFilter + Send + Sync>;
|
||||
pub type BoxNicPacketFilter = Box<dyn NicPacketFilter + Send + Sync>;
|
||||
@@ -0,0 +1,420 @@
|
||||
use std::{
|
||||
collections::BTreeSet,
|
||||
sync::{Arc, Weak},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use crossbeam::atomic::AtomicCell;
|
||||
use futures::Future;
|
||||
use std::sync::RwLock;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::task::JoinSet;
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
peers::{
|
||||
peer_rpc::PeerRpcManager,
|
||||
route::{RouteCostCalculator, RouteCostCalculatorInterface},
|
||||
},
|
||||
proto::{
|
||||
core_peer::peer::Route as CoreRoute,
|
||||
peer_rpc::{
|
||||
GetGlobalPeerMapRequest, GetGlobalPeerMapResponse, GlobalPeerMap, PeerCenterRpc,
|
||||
PeerCenterRpcClientFactory, PeerCenterRpcServer, PeerInfoForGlobalMap,
|
||||
ReportPeersRequest, ReportPeersResponse,
|
||||
},
|
||||
rpc_types::{self, controller::BaseController},
|
||||
},
|
||||
};
|
||||
|
||||
use super::{Digest, Error, server::PeerCenterServer};
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(&, Arc, Box)]
|
||||
pub trait PeerCenterPeerManagerTrait: Send + Sync + 'static {
|
||||
async fn list_peers(&self) -> PeerInfoForGlobalMap;
|
||||
fn my_peer_id(&self) -> PeerId;
|
||||
fn network_name(&self) -> String;
|
||||
fn get_rpc_mgr(&self) -> Weak<PeerRpcManager>;
|
||||
async fn list_routes(&self) -> Vec<CoreRoute>;
|
||||
}
|
||||
|
||||
struct PeerCenterBase {
|
||||
peer_mgr: Arc<dyn PeerCenterPeerManagerTrait>,
|
||||
my_peer_id: PeerId,
|
||||
tasks: Mutex<JoinSet<()>>,
|
||||
lock: Arc<Mutex<()>>,
|
||||
}
|
||||
|
||||
struct PeridicJobCtx<T> {
|
||||
my_peer_id: PeerId,
|
||||
center_peer: AtomicCell<PeerId>,
|
||||
job_ctx: T,
|
||||
}
|
||||
|
||||
impl PeerCenterBase {
|
||||
pub async fn init(&self) -> Result<(), Error> {
|
||||
let Some(rpc_mgr) = self.peer_mgr.get_rpc_mgr().upgrade() else {
|
||||
return Err(Error::Shutdown);
|
||||
};
|
||||
rpc_mgr.rpc_server().registry().register(
|
||||
PeerCenterRpcServer::new(PeerCenterServer::new()),
|
||||
&self.peer_mgr.network_name(),
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn select_center_peer(peer_mgr: &dyn PeerCenterPeerManagerTrait) -> Option<PeerId> {
|
||||
let peers = peer_mgr.list_routes().await;
|
||||
if peers.is_empty() {
|
||||
return None;
|
||||
}
|
||||
// find peer with alphabetical smallest id.
|
||||
let mut min_peer = peer_mgr.my_peer_id();
|
||||
for peer in peers
|
||||
.iter()
|
||||
.filter(|r| r.feature_flag.map(|r| !r.is_public_server).unwrap_or(true))
|
||||
{
|
||||
let peer_id = peer.peer_id;
|
||||
if peer_id < min_peer {
|
||||
min_peer = peer_id;
|
||||
}
|
||||
}
|
||||
Some(min_peer)
|
||||
}
|
||||
|
||||
async fn init_periodic_job<
|
||||
T: Send + Sync + 'static + Clone,
|
||||
Fut: Future<Output = Result<u32, rpc_types::error::Error>> + Send + 'static,
|
||||
>(
|
||||
&self,
|
||||
job_ctx: T,
|
||||
job_fn: impl Fn(
|
||||
Box<dyn PeerCenterRpc<Controller = BaseController> + Send>,
|
||||
Arc<PeridicJobCtx<T>>,
|
||||
) -> Fut
|
||||
+ Send
|
||||
+ Sync
|
||||
+ 'static,
|
||||
) {
|
||||
let my_peer_id = self.my_peer_id;
|
||||
let peer_mgr = self.peer_mgr.clone();
|
||||
let lock = self.lock.clone();
|
||||
self.tasks.lock().await.spawn(
|
||||
async move {
|
||||
let ctx = Arc::new(PeridicJobCtx {
|
||||
my_peer_id,
|
||||
center_peer: AtomicCell::new(PeerId::default()),
|
||||
job_ctx,
|
||||
});
|
||||
loop {
|
||||
let Some(center_peer) = Self::select_center_peer(&peer_mgr).await else {
|
||||
tracing::trace!("no center peer found, sleep 1 second");
|
||||
crate::foundation::time::sleep(Duration::from_secs(1)).await;
|
||||
continue;
|
||||
};
|
||||
let Some(rpc_mgr) = peer_mgr.get_rpc_mgr().upgrade() else {
|
||||
tracing::error!("rpc manager is shutdown, exit periodic job");
|
||||
return;
|
||||
};
|
||||
|
||||
ctx.center_peer.store(center_peer);
|
||||
tracing::trace!(?center_peer, "run periodic job");
|
||||
let _g = lock.lock().await;
|
||||
let stub = rpc_mgr
|
||||
.rpc_client()
|
||||
.scoped_client::<PeerCenterRpcClientFactory<BaseController>>(
|
||||
my_peer_id,
|
||||
center_peer,
|
||||
peer_mgr.network_name(),
|
||||
);
|
||||
let ret = job_fn(stub, ctx.clone()).await;
|
||||
drop(_g);
|
||||
|
||||
let Ok(sleep_time_ms) = ret else {
|
||||
tracing::error!("periodic job to center server rpc failed: {:?}", ret);
|
||||
crate::foundation::time::sleep(Duration::from_secs(3)).await;
|
||||
continue;
|
||||
};
|
||||
|
||||
if sleep_time_ms > 0 {
|
||||
crate::foundation::time::sleep(Duration::from_millis(sleep_time_ms as u64))
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
.instrument(tracing::info_span!("periodic_job", ?my_peer_id)),
|
||||
);
|
||||
}
|
||||
|
||||
pub fn new(peer_mgr: Arc<dyn PeerCenterPeerManagerTrait>) -> Self {
|
||||
let my_peer_id = peer_mgr.my_peer_id();
|
||||
PeerCenterBase {
|
||||
peer_mgr,
|
||||
my_peer_id,
|
||||
tasks: Mutex::new(JoinSet::new()),
|
||||
lock: Arc::new(Mutex::new(())),
|
||||
}
|
||||
}
|
||||
|
||||
async fn stop(&self) {
|
||||
let mut tasks = self.tasks.lock().await;
|
||||
tasks.abort_all();
|
||||
while tasks.join_next().await.is_some() {}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct PeerCenterInstanceService {
|
||||
global_peer_map: Arc<RwLock<GlobalPeerMap>>,
|
||||
global_peer_map_digest: Arc<AtomicCell<Digest>>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PeerCenterRpc for PeerCenterInstanceService {
|
||||
type Controller = BaseController;
|
||||
|
||||
async fn get_global_peer_map(
|
||||
&self,
|
||||
_: BaseController,
|
||||
_: GetGlobalPeerMapRequest,
|
||||
) -> Result<GetGlobalPeerMapResponse, rpc_types::error::Error> {
|
||||
let global_peer_map = self.global_peer_map.read().unwrap();
|
||||
Ok(GetGlobalPeerMapResponse {
|
||||
global_peer_map: global_peer_map.map.clone(),
|
||||
digest: Some(self.global_peer_map_digest.load()),
|
||||
})
|
||||
}
|
||||
|
||||
async fn report_peers(
|
||||
&self,
|
||||
_: BaseController,
|
||||
_req: ReportPeersRequest,
|
||||
) -> Result<ReportPeersResponse, rpc_types::error::Error> {
|
||||
Err(anyhow::anyhow!("not implemented").into())
|
||||
}
|
||||
}
|
||||
|
||||
pub struct PeerCenterInstance {
|
||||
peer_mgr: Arc<dyn PeerCenterPeerManagerTrait>,
|
||||
|
||||
client: Arc<PeerCenterBase>,
|
||||
global_peer_map: Arc<RwLock<GlobalPeerMap>>,
|
||||
global_peer_map_digest: Arc<AtomicCell<Digest>>,
|
||||
global_peer_map_update_time: Arc<AtomicCell<Instant>>,
|
||||
}
|
||||
|
||||
impl PeerCenterInstance {
|
||||
pub fn new(peer_mgr: Arc<dyn PeerCenterPeerManagerTrait>) -> Self {
|
||||
PeerCenterInstance {
|
||||
peer_mgr: peer_mgr.clone(),
|
||||
client: Arc::new(PeerCenterBase::new(peer_mgr.clone())),
|
||||
global_peer_map: Arc::new(RwLock::new(GlobalPeerMap::default())),
|
||||
global_peer_map_digest: Arc::new(AtomicCell::new(Digest::default())),
|
||||
global_peer_map_update_time: Arc::new(AtomicCell::new(Instant::now())),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn global_peer_map_snapshot(&self) -> GetGlobalPeerMapResponse {
|
||||
GetGlobalPeerMapResponse {
|
||||
global_peer_map: self.global_peer_map.read().unwrap().map.clone(),
|
||||
digest: Some(self.global_peer_map_digest.load()),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn init(&self) {
|
||||
self.client.init().await.unwrap();
|
||||
self.init_get_global_info_job().await;
|
||||
self.init_report_peers_job().await;
|
||||
}
|
||||
|
||||
pub async fn stop(&self) {
|
||||
self.client.stop().await;
|
||||
}
|
||||
|
||||
async fn init_get_global_info_job(&self) {
|
||||
struct Ctx {
|
||||
global_peer_map: Arc<RwLock<GlobalPeerMap>>,
|
||||
global_peer_map_digest: Arc<AtomicCell<Digest>>,
|
||||
global_peer_map_update_time: Arc<AtomicCell<Instant>>,
|
||||
}
|
||||
|
||||
let ctx = Arc::new(Ctx {
|
||||
global_peer_map: self.global_peer_map.clone(),
|
||||
global_peer_map_digest: self.global_peer_map_digest.clone(),
|
||||
global_peer_map_update_time: self.global_peer_map_update_time.clone(),
|
||||
});
|
||||
|
||||
self.client
|
||||
.init_periodic_job(ctx, |client, ctx| async move {
|
||||
if ctx
|
||||
.job_ctx
|
||||
.global_peer_map_update_time
|
||||
.load()
|
||||
.elapsed()
|
||||
.as_secs()
|
||||
> 120
|
||||
{
|
||||
ctx.job_ctx.global_peer_map_digest.store(Digest::default());
|
||||
}
|
||||
|
||||
let ret = client
|
||||
.get_global_peer_map(
|
||||
BaseController::default(),
|
||||
GetGlobalPeerMapRequest {
|
||||
digest: ctx.job_ctx.global_peer_map_digest.load(),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
let Ok(resp) = ret else {
|
||||
tracing::error!(
|
||||
"get global info from center server got error result: {:?}",
|
||||
ret
|
||||
);
|
||||
return Ok(10000);
|
||||
};
|
||||
|
||||
if resp == GetGlobalPeerMapResponse::default() {
|
||||
// digest match, no need to update
|
||||
return Ok(15000);
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
"get global info from center server: {:?}, digest: {:?}",
|
||||
resp.global_peer_map,
|
||||
resp.digest
|
||||
);
|
||||
|
||||
*ctx.job_ctx.global_peer_map.write().unwrap() = GlobalPeerMap {
|
||||
map: resp.global_peer_map,
|
||||
};
|
||||
ctx.job_ctx
|
||||
.global_peer_map_digest
|
||||
.store(resp.digest.unwrap_or_default());
|
||||
ctx.job_ctx
|
||||
.global_peer_map_update_time
|
||||
.store(Instant::now());
|
||||
|
||||
Ok(15000)
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
async fn init_report_peers_job(&self) {
|
||||
struct Ctx {
|
||||
peer_mgr: Arc<dyn PeerCenterPeerManagerTrait>,
|
||||
last_report_peers: Mutex<BTreeSet<PeerId>>,
|
||||
|
||||
last_center_peer: AtomicCell<PeerId>,
|
||||
last_report_time: AtomicCell<Instant>,
|
||||
}
|
||||
let ctx = Arc::new(Ctx {
|
||||
peer_mgr: self.peer_mgr.clone(),
|
||||
last_report_peers: Mutex::new(BTreeSet::new()),
|
||||
last_center_peer: AtomicCell::new(PeerId::default()),
|
||||
last_report_time: AtomicCell::new(Instant::now()),
|
||||
});
|
||||
|
||||
self.client
|
||||
.init_periodic_job(ctx, |client, ctx| async move {
|
||||
let my_node_id = ctx.my_peer_id;
|
||||
let peers = ctx.job_ctx.peer_mgr.list_peers().await;
|
||||
let peer_list = peers.direct_peers.keys().copied().collect();
|
||||
let job_ctx = &ctx.job_ctx;
|
||||
|
||||
// only report when:
|
||||
// 1. center peer changed
|
||||
// 2. last report time is more than 60 seconds
|
||||
// 3. peers changed
|
||||
if ctx.center_peer.load() == ctx.job_ctx.last_center_peer.load()
|
||||
&& job_ctx.last_report_time.load().elapsed().as_secs() < 60
|
||||
&& *job_ctx.last_report_peers.lock().await == peer_list
|
||||
{
|
||||
return Ok(5000);
|
||||
}
|
||||
|
||||
let ret = client
|
||||
.report_peers(
|
||||
BaseController::default(),
|
||||
ReportPeersRequest {
|
||||
my_peer_id: my_node_id,
|
||||
peer_infos: Some(peers),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
if ret.is_ok() {
|
||||
ctx.job_ctx.last_center_peer.store(ctx.center_peer.load());
|
||||
*ctx.job_ctx.last_report_peers.lock().await = peer_list;
|
||||
ctx.job_ctx.last_report_time.store(Instant::now());
|
||||
} else {
|
||||
tracing::error!("report peers to center server got error result: {:?}", ret);
|
||||
}
|
||||
|
||||
Ok(5000)
|
||||
})
|
||||
.await;
|
||||
}
|
||||
|
||||
pub fn get_rpc_service(&self) -> PeerCenterInstanceService {
|
||||
PeerCenterInstanceService {
|
||||
global_peer_map: self.global_peer_map.clone(),
|
||||
global_peer_map_digest: self.global_peer_map_digest.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_cost_calculator(&self) -> RouteCostCalculator {
|
||||
struct RouteCostCalculatorImpl {
|
||||
global_peer_map: Arc<RwLock<GlobalPeerMap>>,
|
||||
|
||||
global_peer_map_clone: GlobalPeerMap,
|
||||
|
||||
last_update_time: AtomicCell<Instant>,
|
||||
global_peer_map_update_time: Arc<AtomicCell<Instant>>,
|
||||
}
|
||||
|
||||
impl RouteCostCalculatorImpl {
|
||||
fn directed_cost(&self, src: PeerId, dst: PeerId) -> Option<i32> {
|
||||
self.global_peer_map_clone
|
||||
.map
|
||||
.get(&src)
|
||||
.and_then(|src_peer_info| src_peer_info.direct_peers.get(&dst))
|
||||
.map(|info| info.latency_ms)
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteCostCalculatorInterface for RouteCostCalculatorImpl {
|
||||
fn calculate_cost(&self, src: PeerId, dst: PeerId) -> i32 {
|
||||
if let Some(cost) = self.directed_cost(src, dst) {
|
||||
return cost;
|
||||
}
|
||||
self.directed_cost(dst, src).unwrap_or(500)
|
||||
}
|
||||
|
||||
fn begin_update(&mut self) {
|
||||
let global_peer_map = self.global_peer_map.read().unwrap();
|
||||
self.global_peer_map_clone = global_peer_map.clone();
|
||||
}
|
||||
|
||||
fn end_update(&mut self) {
|
||||
self.last_update_time
|
||||
.store(self.global_peer_map_update_time.load());
|
||||
}
|
||||
|
||||
fn need_update(&self) -> bool {
|
||||
self.last_update_time.load() < self.global_peer_map_update_time.load()
|
||||
}
|
||||
}
|
||||
|
||||
Box::new(RouteCostCalculatorImpl {
|
||||
global_peer_map: self.global_peer_map.clone(),
|
||||
global_peer_map_clone: GlobalPeerMap::default(),
|
||||
last_update_time: AtomicCell::new(
|
||||
self.global_peer_map_update_time.load() - Duration::from_secs(1),
|
||||
),
|
||||
global_peer_map_update_time: self.global_peer_map_update_time.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
// peer_center is used to collect peer info into one peer node.
|
||||
// the center node is selected with the following rules:
|
||||
// 1. has smallest peer id
|
||||
// 2. TODO: has allow_to_be_center peer feature
|
||||
// peer center is not guaranteed to be stable and can be changed when peer enter or leave.
|
||||
// it's used to reduce the cost to exchange infos between peers.
|
||||
|
||||
pub mod instance;
|
||||
mod server;
|
||||
|
||||
#[derive(thiserror::Error, Debug, serde::Deserialize, serde::Serialize)]
|
||||
pub enum Error {
|
||||
#[error("Digest not match, need provide full peer info to center server.")]
|
||||
DigestMismatch,
|
||||
#[error("Not center server")]
|
||||
NotCenterServer,
|
||||
#[error("Instance shutdown")]
|
||||
Shutdown,
|
||||
}
|
||||
|
||||
pub type Digest = u64;
|
||||
@@ -0,0 +1,239 @@
|
||||
use std::{
|
||||
collections::BinaryHeap,
|
||||
hash::{Hash, Hasher},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use crossbeam::atomic::AtomicCell;
|
||||
use dashmap::DashMap;
|
||||
use tokio::task::JoinSet;
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
proto::{
|
||||
peer_rpc::{
|
||||
DirectConnectedPeerInfo, GetGlobalPeerMapRequest, GetGlobalPeerMapResponse,
|
||||
GlobalPeerMap, PeerCenterRpc, PeerInfoForGlobalMap, ReportPeersRequest,
|
||||
ReportPeersResponse,
|
||||
},
|
||||
rpc_types::{self, controller::BaseController},
|
||||
},
|
||||
};
|
||||
|
||||
use super::Digest;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, PartialOrd, Ord, Eq, Hash)]
|
||||
pub(crate) struct SrcDstPeerPair {
|
||||
src: PeerId,
|
||||
dst: PeerId,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct PeerCenterInfoEntry {
|
||||
info: DirectConnectedPeerInfo,
|
||||
update_time: std::time::Instant,
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct PeerCenterServerData {
|
||||
global_peer_map: DashMap<SrcDstPeerPair, PeerCenterInfoEntry>,
|
||||
peer_report_time: DashMap<PeerId, std::time::Instant>,
|
||||
digest: AtomicCell<Digest>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct PeerCenterServer {
|
||||
data: Arc<PeerCenterServerData>,
|
||||
_tasks: Arc<JoinSet<()>>,
|
||||
}
|
||||
|
||||
impl PeerCenterServer {
|
||||
pub fn new() -> Self {
|
||||
let data = Arc::new(PeerCenterServerData::default());
|
||||
let weak_data = Arc::downgrade(&data);
|
||||
let mut tasks = JoinSet::new();
|
||||
tasks.spawn(async move {
|
||||
loop {
|
||||
crate::foundation::time::sleep(std::time::Duration::from_secs(10)).await;
|
||||
let Some(data) = weak_data.upgrade() else {
|
||||
break;
|
||||
};
|
||||
PeerCenterServer::clean_outdated_peer_data(&data).await;
|
||||
}
|
||||
});
|
||||
|
||||
PeerCenterServer {
|
||||
data,
|
||||
_tasks: Arc::new(tasks),
|
||||
}
|
||||
}
|
||||
|
||||
async fn clean_outdated_peer_data(data: &PeerCenterServerData) {
|
||||
data.peer_report_time.retain(|_, v| {
|
||||
std::time::Instant::now().duration_since(*v) < std::time::Duration::from_secs(180)
|
||||
});
|
||||
data.global_peer_map.retain(|_, v| {
|
||||
std::time::Instant::now().duration_since(v.update_time)
|
||||
< std::time::Duration::from_secs(180)
|
||||
});
|
||||
}
|
||||
|
||||
fn calc_global_digest_data(data: &PeerCenterServerData) -> Digest {
|
||||
let mut hasher = std::collections::hash_map::DefaultHasher::new();
|
||||
data.global_peer_map
|
||||
.iter()
|
||||
.map(|v| v.key().clone())
|
||||
.collect::<BinaryHeap<_>>()
|
||||
.into_sorted_vec()
|
||||
.into_iter()
|
||||
.for_each(|v| v.hash(&mut hasher));
|
||||
hasher.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PeerCenterRpc for PeerCenterServer {
|
||||
type Controller = BaseController;
|
||||
|
||||
#[tracing::instrument()]
|
||||
async fn report_peers(
|
||||
&self,
|
||||
_: BaseController,
|
||||
req: ReportPeersRequest,
|
||||
) -> Result<ReportPeersResponse, rpc_types::error::Error> {
|
||||
let my_peer_id = req.my_peer_id;
|
||||
let peers = req.peer_infos.unwrap_or_default();
|
||||
|
||||
tracing::debug!("receive report_peers");
|
||||
|
||||
let data = &self.data;
|
||||
data.peer_report_time
|
||||
.insert(my_peer_id, std::time::Instant::now());
|
||||
|
||||
for (peer_id, peer_info) in peers.direct_peers {
|
||||
let pair = SrcDstPeerPair {
|
||||
src: my_peer_id,
|
||||
dst: peer_id,
|
||||
};
|
||||
let entry = PeerCenterInfoEntry {
|
||||
info: peer_info,
|
||||
update_time: std::time::Instant::now(),
|
||||
};
|
||||
data.global_peer_map.insert(pair, entry);
|
||||
}
|
||||
|
||||
data.digest
|
||||
.store(PeerCenterServer::calc_global_digest_data(data));
|
||||
|
||||
Ok(ReportPeersResponse::default())
|
||||
}
|
||||
|
||||
#[tracing::instrument()]
|
||||
async fn get_global_peer_map(
|
||||
&self,
|
||||
_: BaseController,
|
||||
req: GetGlobalPeerMapRequest,
|
||||
) -> Result<GetGlobalPeerMapResponse, rpc_types::error::Error> {
|
||||
let digest = req.digest;
|
||||
|
||||
let data = &self.data;
|
||||
if digest == data.digest.load() && digest != 0 {
|
||||
return Ok(GetGlobalPeerMapResponse::default());
|
||||
}
|
||||
|
||||
let mut global_peer_map = GlobalPeerMap::default();
|
||||
for item in data.global_peer_map.iter() {
|
||||
let (pair, entry) = item.pair();
|
||||
global_peer_map
|
||||
.map
|
||||
.entry(pair.src)
|
||||
.or_insert_with(|| PeerInfoForGlobalMap {
|
||||
direct_peers: Default::default(),
|
||||
})
|
||||
.direct_peers
|
||||
.insert(pair.dst, entry.info);
|
||||
}
|
||||
|
||||
Ok(GetGlobalPeerMapResponse {
|
||||
global_peer_map: global_peer_map.map,
|
||||
digest: Some(data.digest.load()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn server_clones_share_instance_data() {
|
||||
let server = PeerCenterServer::new();
|
||||
let server_clone = server.clone();
|
||||
|
||||
let mut peers = PeerInfoForGlobalMap::default();
|
||||
peers
|
||||
.direct_peers
|
||||
.insert(100, DirectConnectedPeerInfo { latency_ms: 3 });
|
||||
|
||||
server
|
||||
.report_peers(
|
||||
BaseController::default(),
|
||||
ReportPeersRequest {
|
||||
my_peer_id: 99,
|
||||
peer_infos: Some(peers),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let resp = server_clone
|
||||
.get_global_peer_map(
|
||||
BaseController::default(),
|
||||
GetGlobalPeerMapRequest { digest: 0 },
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(1, resp.global_peer_map.len());
|
||||
assert!(resp.global_peer_map[&99].direct_peers.contains_key(&100));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn independent_server_instances_do_not_share_data() {
|
||||
let server_a = PeerCenterServer::new();
|
||||
let server_b = PeerCenterServer::new();
|
||||
|
||||
let mut peers = PeerInfoForGlobalMap::default();
|
||||
peers
|
||||
.direct_peers
|
||||
.insert(101, DirectConnectedPeerInfo { latency_ms: 5 });
|
||||
|
||||
server_a
|
||||
.report_peers(
|
||||
BaseController::default(),
|
||||
ReportPeersRequest {
|
||||
my_peer_id: 100,
|
||||
peer_infos: Some(peers),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let resp_a = server_a
|
||||
.get_global_peer_map(
|
||||
BaseController::default(),
|
||||
GetGlobalPeerMapRequest { digest: 0 },
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(1, resp_a.global_peer_map.len());
|
||||
|
||||
let resp_b = server_b
|
||||
.get_global_peer_map(
|
||||
BaseController::default(),
|
||||
GetGlobalPeerMapRequest { digest: 0 },
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(resp_b.global_peer_map.is_empty());
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,113 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use futures::{SinkExt as _, StreamExt};
|
||||
use tokio::task::JoinSet;
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
foundation::stats::{ArcRpcMetrics, RpcMetricsProvider},
|
||||
packet::ZCPacket,
|
||||
rpc::{self, bidirect::BidirectRpcManager},
|
||||
};
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(Arc)]
|
||||
pub trait PeerRpcManagerTransport: Send + Sync + 'static {
|
||||
fn my_peer_id(&self) -> PeerId;
|
||||
async fn send(&self, msg: ZCPacket, dst_peer_id: PeerId) -> anyhow::Result<()>;
|
||||
async fn recv(&self) -> anyhow::Result<ZCPacket>;
|
||||
}
|
||||
|
||||
pub struct PeerRpcManager {
|
||||
tspt: Arc<Box<dyn PeerRpcManagerTransport>>,
|
||||
bidirect_rpc: BidirectRpcManager,
|
||||
tasks: Mutex<JoinSet<()>>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PeerRpcManager {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("PeerRpcManager")
|
||||
.field("node_id", &self.tspt.my_peer_id())
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerRpcManager {
|
||||
pub fn new(tspt: impl PeerRpcManagerTransport) -> Self {
|
||||
Self {
|
||||
tspt: Arc::new(Box::new(tspt)),
|
||||
bidirect_rpc: BidirectRpcManager::new(),
|
||||
tasks: Mutex::new(JoinSet::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_with_stats_manager<T>(tspt: impl PeerRpcManagerTransport, stats_manager: T) -> Self
|
||||
where
|
||||
T: Clone + RpcMetricsProvider,
|
||||
{
|
||||
Self {
|
||||
tspt: Arc::new(Box::new(tspt)),
|
||||
bidirect_rpc: BidirectRpcManager::new_with_stats_manager(stats_manager),
|
||||
tasks: Mutex::new(JoinSet::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_with_metrics(tspt: impl PeerRpcManagerTransport, metrics: ArcRpcMetrics) -> Self {
|
||||
Self {
|
||||
tspt: Arc::new(Box::new(tspt)),
|
||||
bidirect_rpc: BidirectRpcManager::new_with_metrics(metrics),
|
||||
tasks: Mutex::new(JoinSet::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn run(&self) {
|
||||
let ret = self.bidirect_rpc.run_and_create_tunnel();
|
||||
let (mut rx, mut tx) = ret.split();
|
||||
let tspt = self.tspt.clone();
|
||||
self.tasks.lock().unwrap().spawn(async move {
|
||||
while let Some(Ok(packet)) = rx.next().await {
|
||||
let dst_peer_id = packet.peer_manager_header().unwrap().to_peer_id.into();
|
||||
if let Err(e) = tspt.send(packet, dst_peer_id).await {
|
||||
tracing::error!("send to rpc tspt error: {:?}", e);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let tspt = self.tspt.clone();
|
||||
self.tasks.lock().unwrap().spawn(async move {
|
||||
while let Ok(packet) = tspt.recv().await {
|
||||
if let Err(e) = tx.send(packet).await {
|
||||
tracing::error!("send to rpc tspt error: {:?}", e);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
pub async fn stop(&self) {
|
||||
self.bidirect_rpc.stop().await;
|
||||
let mut tasks = {
|
||||
let mut task_slot = self.tasks.lock().unwrap();
|
||||
std::mem::replace(&mut *task_slot, JoinSet::new())
|
||||
};
|
||||
tasks.abort_all();
|
||||
while tasks.join_next().await.is_some() {}
|
||||
}
|
||||
|
||||
pub fn rpc_client(&self) -> &rpc::client::Client {
|
||||
self.bidirect_rpc.rpc_client()
|
||||
}
|
||||
|
||||
pub fn rpc_server(&self) -> &rpc::server::Server {
|
||||
self.bidirect_rpc.rpc_server()
|
||||
}
|
||||
|
||||
pub fn my_peer_id(&self) -> PeerId {
|
||||
self.tspt.my_peer_id()
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for PeerRpcManager {
|
||||
fn drop(&mut self) {
|
||||
tracing::debug!("PeerRpcManager drop, my_peer_id: {:?}", self.my_peer_id());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,738 @@
|
||||
pub mod provider;
|
||||
pub(crate) mod service;
|
||||
|
||||
pub(crate) use service::PublicIpv6Service;
|
||||
|
||||
use std::{collections::HashSet, net::Ipv6Addr, sync::Arc};
|
||||
|
||||
use cidr::{Ipv6Cidr, Ipv6Inet};
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
config::peers::PublicIpv6ProviderConfig,
|
||||
config::runtime::CoreRuntimeConfigStore,
|
||||
events::{CoreEvent, CoreEventSink},
|
||||
peers::context::PeerPublicIpv6State,
|
||||
};
|
||||
|
||||
impl PublicIpv6ProviderConfig {
|
||||
pub fn validate(self) -> Result<(), PublicIpv6ProviderConfigError> {
|
||||
if !self.provider_enabled {
|
||||
return Ok(());
|
||||
}
|
||||
if !self.provider_supported {
|
||||
return Err(PublicIpv6ProviderConfigError::UnsupportedProvider);
|
||||
}
|
||||
if let Some(prefix) = self.configured_prefix
|
||||
&& !is_global_routable_public_ipv6_prefix(prefix)
|
||||
{
|
||||
return Err(PublicIpv6ProviderConfigError::InvalidPrefix(prefix));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
|
||||
pub enum PublicIpv6ProviderConfigError {
|
||||
#[error(
|
||||
"the provider feature requires Linux; run without --ipv6-public-addr-provider on this node, or move the provider role to a Linux node. client mode (--ipv6-public-addr-auto) works on all platforms"
|
||||
)]
|
||||
UnsupportedProvider,
|
||||
#[error(
|
||||
"the prefix {0} is not a valid global unicast IPv6 prefix; it must be a routable address range, not a private, link-local, or multicast address"
|
||||
)]
|
||||
InvalidPrefix(Ipv6Cidr),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) enum PublicIpv6ProviderResolution {
|
||||
Disabled,
|
||||
Pending(String),
|
||||
Active(Ipv6Cidr),
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_public_ipv6_provider(
|
||||
config: PublicIpv6ProviderConfig,
|
||||
detected_prefix: Result<Option<Ipv6Cidr>, String>,
|
||||
) -> PublicIpv6ProviderResolution {
|
||||
if !config.provider_enabled {
|
||||
return PublicIpv6ProviderResolution::Disabled;
|
||||
}
|
||||
if !config.provider_supported {
|
||||
return PublicIpv6ProviderResolution::Pending(
|
||||
PublicIpv6ProviderConfigError::UnsupportedProvider.to_string(),
|
||||
);
|
||||
}
|
||||
|
||||
if let Some(prefix) = config.configured_prefix {
|
||||
return if is_global_routable_public_ipv6_prefix(prefix) {
|
||||
PublicIpv6ProviderResolution::Active(prefix)
|
||||
} else {
|
||||
PublicIpv6ProviderResolution::Pending(format!(
|
||||
"the configured prefix {prefix} is not a valid global unicast IPv6 prefix"
|
||||
))
|
||||
};
|
||||
}
|
||||
|
||||
match detected_prefix {
|
||||
Ok(Some(prefix)) if is_global_routable_public_ipv6_prefix(prefix) => {
|
||||
PublicIpv6ProviderResolution::Active(prefix)
|
||||
}
|
||||
Ok(Some(prefix)) => PublicIpv6ProviderResolution::Pending(format!(
|
||||
"the detected prefix {prefix} is not a valid global unicast IPv6 prefix"
|
||||
)),
|
||||
Ok(None) => PublicIpv6ProviderResolution::Pending(
|
||||
"no public IPv6 prefix found on this system; set --ipv6-public-addr-prefix manually, or check that your ISP has delegated an IPv6 prefix and a default-from route exists in the kernel routing table".to_owned(),
|
||||
),
|
||||
Err(error) => PublicIpv6ProviderResolution::Pending(error),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_global_routable_public_ipv6_prefix(prefix: Ipv6Cidr) -> bool {
|
||||
let addr = prefix.first_address();
|
||||
!addr.is_loopback()
|
||||
&& !addr.is_multicast()
|
||||
&& !addr.is_unicast_link_local()
|
||||
&& !addr.is_unique_local()
|
||||
&& !addr.is_unspecified()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct PublicIpv6PeerRouteInfo {
|
||||
pub peer_id: PeerId,
|
||||
pub inst_id: Option<uuid::Uuid>,
|
||||
pub is_provider: bool,
|
||||
pub prefix: Option<Ipv6Cidr>,
|
||||
pub lease: Option<Ipv6Inet>,
|
||||
pub reachable: bool,
|
||||
}
|
||||
|
||||
pub(crate) trait PublicIpv6RouteControl: Send + Sync {
|
||||
fn my_peer_id(&self) -> PeerId;
|
||||
fn peer_route_snapshot(&self) -> Vec<PublicIpv6PeerRouteInfo>;
|
||||
fn publish_self_public_ipv6_lease(&self, lease: Option<Ipv6Inet>) -> bool;
|
||||
}
|
||||
|
||||
pub(crate) trait PublicIpv6SyncTrigger: Send + Sync {
|
||||
fn sync_now(&self, reason: &str);
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub trait PublicIpv6Host: Send + Sync {
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet<Ipv6Addr>;
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PublicIpv6Host for () {
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, _prefix: Ipv6Cidr) -> HashSet<Ipv6Addr> {
|
||||
HashSet::new()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(Arc)]
|
||||
pub(crate) trait PublicIpv6Runtime: Send + Sync {
|
||||
fn ipv6_public_addr_auto(&self) -> bool;
|
||||
fn ipv6_public_addr_provider(&self) -> bool;
|
||||
fn instance_id(&self) -> uuid::Uuid;
|
||||
fn network_name(&self) -> String;
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet<Ipv6Addr>;
|
||||
fn public_ipv6_lease_changed(&self, old: Option<Ipv6Inet>, new: Option<Ipv6Inet>);
|
||||
fn public_ipv6_routes_changed(&self, added: Vec<Ipv6Inet>, removed: Vec<Ipv6Inet>);
|
||||
}
|
||||
|
||||
pub struct CorePublicIpv6Runtime {
|
||||
config: CoreRuntimeConfigStore,
|
||||
host: Arc<dyn PublicIpv6Host>,
|
||||
events: Arc<dyn CoreEventSink>,
|
||||
provider_prefix: std::sync::Mutex<Option<Ipv6Cidr>>,
|
||||
lease: std::sync::Mutex<Option<Ipv6Inet>>,
|
||||
}
|
||||
|
||||
impl CorePublicIpv6Runtime {
|
||||
pub fn new(
|
||||
config: CoreRuntimeConfigStore,
|
||||
host: Arc<dyn PublicIpv6Host>,
|
||||
events: Arc<dyn CoreEventSink>,
|
||||
) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
config,
|
||||
host,
|
||||
events,
|
||||
provider_prefix: std::sync::Mutex::new(None),
|
||||
lease: std::sync::Mutex::new(None),
|
||||
})
|
||||
}
|
||||
|
||||
pub fn set_provider_prefix(&self, prefix: Option<Ipv6Cidr>) -> bool {
|
||||
let mut current = self.provider_prefix.lock().unwrap();
|
||||
if *current == prefix {
|
||||
return false;
|
||||
}
|
||||
*current = prefix;
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerPublicIpv6State for CorePublicIpv6Runtime {
|
||||
fn public_ipv6_lease_contains(&self, ip: &Ipv6Addr) -> bool {
|
||||
self.lease
|
||||
.lock()
|
||||
.unwrap()
|
||||
.is_some_and(|lease| lease.address() == *ip)
|
||||
}
|
||||
|
||||
fn public_ipv6_provider_enabled(&self) -> bool {
|
||||
self.provider_prefix.lock().unwrap().is_some()
|
||||
}
|
||||
|
||||
fn advertised_ipv6_public_addr_prefix(&self) -> Option<Ipv6Cidr> {
|
||||
*self.provider_prefix.lock().unwrap()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PublicIpv6Runtime for CorePublicIpv6Runtime {
|
||||
fn ipv6_public_addr_auto(&self) -> bool {
|
||||
self.config.snapshot().services.public_ipv6_auto
|
||||
}
|
||||
|
||||
fn ipv6_public_addr_provider(&self) -> bool {
|
||||
self.config
|
||||
.snapshot()
|
||||
.services
|
||||
.public_ipv6_provider
|
||||
.provider_enabled
|
||||
}
|
||||
|
||||
fn instance_id(&self) -> uuid::Uuid {
|
||||
self.config
|
||||
.snapshot()
|
||||
.peer
|
||||
.runtime
|
||||
.core
|
||||
.node
|
||||
.instance_id
|
||||
.map(uuid::Uuid::from_bytes)
|
||||
.expect("core peer identity must be finalized before public IPv6 starts")
|
||||
}
|
||||
|
||||
fn network_name(&self) -> String {
|
||||
self.config
|
||||
.snapshot()
|
||||
.peer
|
||||
.runtime
|
||||
.network_identity
|
||||
.network_name
|
||||
.clone()
|
||||
}
|
||||
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet<Ipv6Addr> {
|
||||
self.host.collect_reserved_public_ipv6_addrs(prefix).await
|
||||
}
|
||||
|
||||
fn public_ipv6_lease_changed(&self, old: Option<Ipv6Inet>, new: Option<Ipv6Inet>) {
|
||||
*self.lease.lock().unwrap() = new;
|
||||
self.events
|
||||
.emit(CoreEvent::PublicIpv6LeaseChanged { old, new });
|
||||
}
|
||||
|
||||
fn public_ipv6_routes_changed(&self, added: Vec<Ipv6Inet>, removed: Vec<Ipv6Inet>) {
|
||||
self.events
|
||||
.emit(CoreEvent::PublicIpv6RoutesChanged { added, removed });
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct DisabledPublicIpv6Runtime {
|
||||
instance_id: uuid::Uuid,
|
||||
network_name: String,
|
||||
}
|
||||
|
||||
impl DisabledPublicIpv6Runtime {
|
||||
pub(super) fn new(instance_id: uuid::Uuid, network_name: String) -> Self {
|
||||
Self {
|
||||
instance_id,
|
||||
network_name,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PublicIpv6Runtime for DisabledPublicIpv6Runtime {
|
||||
fn ipv6_public_addr_auto(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn ipv6_public_addr_provider(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn instance_id(&self) -> uuid::Uuid {
|
||||
self.instance_id
|
||||
}
|
||||
|
||||
fn network_name(&self) -> String {
|
||||
self.network_name.clone()
|
||||
}
|
||||
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, _prefix: Ipv6Cidr) -> HashSet<Ipv6Addr> {
|
||||
HashSet::new()
|
||||
}
|
||||
|
||||
fn public_ipv6_lease_changed(&self, _old: Option<Ipv6Inet>, _new: Option<Ipv6Inet>) {}
|
||||
|
||||
fn public_ipv6_routes_changed(&self, _added: Vec<Ipv6Inet>, _removed: Vec<Ipv6Inet>) {}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::net::Ipv6Addr;
|
||||
use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use cidr::{Ipv6Cidr, Ipv6Inet};
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore},
|
||||
events::{CoreEvent, CoreEventSink},
|
||||
peers::{context::PeerPublicIpv6State, peer_rpc::PeerRpcManager},
|
||||
};
|
||||
|
||||
use super::{
|
||||
CorePublicIpv6Runtime, PublicIpv6Host, PublicIpv6PeerRouteInfo, PublicIpv6ProviderConfig,
|
||||
PublicIpv6ProviderConfigError, PublicIpv6ProviderResolution, PublicIpv6RouteControl,
|
||||
PublicIpv6Runtime, PublicIpv6Service, PublicIpv6SyncTrigger, resolve_public_ipv6_provider,
|
||||
service::allocate_public_ipv6_leases,
|
||||
};
|
||||
|
||||
struct TestRouteControl {
|
||||
my_peer_id: PeerId,
|
||||
peers: Mutex<Vec<PublicIpv6PeerRouteInfo>>,
|
||||
}
|
||||
|
||||
impl PublicIpv6RouteControl for TestRouteControl {
|
||||
fn my_peer_id(&self) -> PeerId {
|
||||
self.my_peer_id
|
||||
}
|
||||
|
||||
fn peer_route_snapshot(&self) -> Vec<PublicIpv6PeerRouteInfo> {
|
||||
self.peers.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
fn publish_self_public_ipv6_lease(&self, _lease: Option<Ipv6Inet>) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
struct TestSyncTrigger;
|
||||
|
||||
impl PublicIpv6SyncTrigger for TestSyncTrigger {
|
||||
fn sync_now(&self, _reason: &str) {}
|
||||
}
|
||||
|
||||
struct TestRuntime {
|
||||
auto: bool,
|
||||
provider: bool,
|
||||
inst_id: uuid::Uuid,
|
||||
network_name: String,
|
||||
reserved: Mutex<HashSet<Ipv6Addr>>,
|
||||
lease: Mutex<Option<Ipv6Inet>>,
|
||||
}
|
||||
|
||||
impl TestRuntime {
|
||||
fn new(auto: bool) -> Self {
|
||||
Self {
|
||||
auto,
|
||||
provider: false,
|
||||
inst_id: uuid::Uuid::from_u128(1),
|
||||
network_name: "default".to_string(),
|
||||
reserved: Mutex::new(HashSet::new()),
|
||||
lease: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PublicIpv6Runtime for TestRuntime {
|
||||
fn ipv6_public_addr_auto(&self) -> bool {
|
||||
self.auto
|
||||
}
|
||||
|
||||
fn ipv6_public_addr_provider(&self) -> bool {
|
||||
self.provider
|
||||
}
|
||||
|
||||
fn instance_id(&self) -> uuid::Uuid {
|
||||
self.inst_id
|
||||
}
|
||||
|
||||
fn network_name(&self) -> String {
|
||||
self.network_name.clone()
|
||||
}
|
||||
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet<Ipv6Addr> {
|
||||
self.reserved
|
||||
.lock()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.copied()
|
||||
.filter(|addr| prefix.contains(addr))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn public_ipv6_lease_changed(&self, _old: Option<Ipv6Inet>, new: Option<Ipv6Inet>) {
|
||||
*self.lease.lock().unwrap() = new;
|
||||
}
|
||||
|
||||
fn public_ipv6_routes_changed(&self, _added: Vec<Ipv6Inet>, _removed: Vec<Ipv6Inet>) {}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct RecordingPublicIpv6Host {
|
||||
reserved: Mutex<HashSet<Ipv6Addr>>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct RecordingPublicIpv6Events {
|
||||
leases: Mutex<Vec<(Option<Ipv6Inet>, Option<Ipv6Inet>)>>,
|
||||
route_deltas: Mutex<Vec<(Vec<Ipv6Inet>, Vec<Ipv6Inet>)>>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PublicIpv6Host for RecordingPublicIpv6Host {
|
||||
async fn collect_reserved_public_ipv6_addrs(&self, prefix: Ipv6Cidr) -> HashSet<Ipv6Addr> {
|
||||
self.reserved
|
||||
.lock()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.copied()
|
||||
.filter(|addr| prefix.contains(addr))
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl CoreEventSink for RecordingPublicIpv6Events {
|
||||
fn emit(&self, event: CoreEvent) {
|
||||
match event {
|
||||
CoreEvent::PublicIpv6LeaseChanged { old, new } => {
|
||||
self.leases.lock().unwrap().push((old, new));
|
||||
}
|
||||
CoreEvent::PublicIpv6RoutesChanged { added, removed } => {
|
||||
self.route_deltas.lock().unwrap().push((added, removed));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn core_runtime_owns_public_ipv6_state_and_projects_only_host_effects() {
|
||||
let instance_id = uuid::Uuid::from_u128(42);
|
||||
let mut peer = crate::config::peers::PeerRuntimeSnapshot::default();
|
||||
peer.runtime.core.node.instance_id = Some(*instance_id.as_bytes());
|
||||
peer.runtime.network_identity.network_name = "owned-by-core".to_owned();
|
||||
let config = CoreRuntimeConfigStore::new(
|
||||
CoreRuntimeConfig {
|
||||
public_ipv6_auto: true,
|
||||
public_ipv6_provider: PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: None,
|
||||
provider_supported: true,
|
||||
},
|
||||
..Default::default()
|
||||
},
|
||||
Arc::new(peer),
|
||||
);
|
||||
let host = Arc::new(RecordingPublicIpv6Host::default());
|
||||
let events = Arc::new(RecordingPublicIpv6Events::default());
|
||||
let reserved = "2001:db8::10".parse().unwrap();
|
||||
host.reserved.lock().unwrap().insert(reserved);
|
||||
let runtime = CorePublicIpv6Runtime::new(config.clone(), host.clone(), events.clone());
|
||||
let prefix = "2001:db8::/64".parse().unwrap();
|
||||
let lease = "2001:db8::20/64".parse().unwrap();
|
||||
let route = "2001:db8::30/128".parse().unwrap();
|
||||
|
||||
assert!(runtime.ipv6_public_addr_auto());
|
||||
assert!(runtime.ipv6_public_addr_provider());
|
||||
assert_eq!(runtime.instance_id(), instance_id);
|
||||
assert_eq!(runtime.network_name(), "owned-by-core");
|
||||
assert_eq!(
|
||||
runtime.collect_reserved_public_ipv6_addrs(prefix).await,
|
||||
HashSet::from([reserved])
|
||||
);
|
||||
assert!(runtime.set_provider_prefix(Some(prefix)));
|
||||
assert!(!runtime.set_provider_prefix(Some(prefix)));
|
||||
assert_eq!(runtime.advertised_ipv6_public_addr_prefix(), Some(prefix));
|
||||
|
||||
runtime.public_ipv6_lease_changed(None, Some(lease));
|
||||
assert!(runtime.public_ipv6_lease_contains(&lease.address()));
|
||||
runtime.public_ipv6_routes_changed(vec![route], Vec::new());
|
||||
assert_eq!(
|
||||
events.leases.lock().unwrap().as_slice(),
|
||||
&[(None, Some(lease))]
|
||||
);
|
||||
assert_eq!(
|
||||
events.route_deltas.lock().unwrap().as_slice(),
|
||||
&[(vec![route], Vec::new())]
|
||||
);
|
||||
|
||||
config.update_services(|services| {
|
||||
services.public_ipv6_auto = false;
|
||||
services.public_ipv6_provider.provider_enabled = false;
|
||||
});
|
||||
assert!(!runtime.ipv6_public_addr_auto());
|
||||
assert!(!runtime.ipv6_public_addr_provider());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn public_ipv6_lease_allocator_keeps_stable_addresses() {
|
||||
let prefix = "2001:db8::/124".parse::<Ipv6Cidr>().unwrap();
|
||||
let first = uuid::Uuid::from_u128(1);
|
||||
let second = uuid::Uuid::from_u128(2);
|
||||
|
||||
let leases =
|
||||
allocate_public_ipv6_leases(prefix, &[first, second], &HashSet::new(), &HashMap::new());
|
||||
assert_eq!(leases.len(), 2);
|
||||
assert_ne!(leases[0].addr, leases[1].addr);
|
||||
|
||||
let initial_map = HashMap::from([(first, leases[0].addr)]);
|
||||
let next = allocate_public_ipv6_leases(prefix, &[first], &HashSet::new(), &initial_map);
|
||||
assert_eq!(next.len(), 1);
|
||||
assert_eq!(next[0].addr, leases[0].addr);
|
||||
assert!(next[0].reused);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn public_ipv6_provider_prefers_smallest_instance_id() {
|
||||
let info_a = PublicIpv6PeerRouteInfo {
|
||||
peer_id: 2,
|
||||
inst_id: Some(uuid::Uuid::from_u128(2)),
|
||||
is_provider: true,
|
||||
prefix: Some("2001:db8:1::/120".parse().unwrap()),
|
||||
lease: None,
|
||||
reachable: true,
|
||||
};
|
||||
let info_b = PublicIpv6PeerRouteInfo {
|
||||
peer_id: 1,
|
||||
inst_id: Some(uuid::Uuid::from_u128(1)),
|
||||
is_provider: true,
|
||||
prefix: Some("2001:db8:2::/120".parse().unwrap()),
|
||||
lease: None,
|
||||
reachable: true,
|
||||
};
|
||||
|
||||
let selected =
|
||||
PublicIpv6Service::selected_provider_from_snapshot(&[info_a, info_b]).unwrap();
|
||||
assert_eq!(selected.peer_id, 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn public_ipv6_provider_prefers_reachable_provider() {
|
||||
let unreachable_lower_id = PublicIpv6PeerRouteInfo {
|
||||
peer_id: 1,
|
||||
inst_id: Some(uuid::Uuid::from_u128(1)),
|
||||
is_provider: true,
|
||||
prefix: Some("2001:db8:1::/120".parse().unwrap()),
|
||||
lease: None,
|
||||
reachable: false,
|
||||
};
|
||||
let reachable_higher_id = PublicIpv6PeerRouteInfo {
|
||||
peer_id: 2,
|
||||
inst_id: Some(uuid::Uuid::from_u128(2)),
|
||||
is_provider: true,
|
||||
prefix: Some("2001:db8:2::/120".parse().unwrap()),
|
||||
lease: None,
|
||||
reachable: true,
|
||||
};
|
||||
|
||||
let selected = PublicIpv6Service::selected_provider_from_snapshot(&[
|
||||
unreachable_lower_id,
|
||||
reachable_higher_id,
|
||||
])
|
||||
.unwrap();
|
||||
assert_eq!(selected.peer_id, 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn public_ipv6_lease_allocator_stops_when_only_network_offset_is_left() {
|
||||
let prefix = "2001:db8::/126".parse::<Ipv6Cidr>().unwrap();
|
||||
let network = prefix.first_address();
|
||||
let reserved = HashSet::from([
|
||||
Ipv6Addr::from(u128::from(network) + 1),
|
||||
Ipv6Addr::from(u128::from(network) + 2),
|
||||
Ipv6Addr::from(u128::from(network) + 3),
|
||||
]);
|
||||
|
||||
let leases = allocate_public_ipv6_leases(
|
||||
prefix,
|
||||
&[uuid::Uuid::from_u128(42)],
|
||||
&reserved,
|
||||
&HashMap::new(),
|
||||
);
|
||||
|
||||
assert!(leases.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reconcile_runtime_clears_public_ipv6_lease_when_auto_is_disabled() {
|
||||
let stale_addr = "2001:db8::123/64".parse().unwrap();
|
||||
let runtime = Arc::new(TestRuntime::new(false));
|
||||
*runtime.lease.lock().unwrap() = Some(stale_addr);
|
||||
|
||||
let service = Arc::new(PublicIpv6Service::new(
|
||||
runtime.clone(),
|
||||
std::sync::Weak::<PeerRpcManager>::new(),
|
||||
Arc::new(TestRouteControl {
|
||||
my_peer_id: 1,
|
||||
peers: Mutex::new(Vec::new()),
|
||||
}),
|
||||
Arc::new(TestSyncTrigger),
|
||||
));
|
||||
*service.my_addr_cache.lock().unwrap() = Some(stale_addr);
|
||||
|
||||
service.reconcile_runtime_from_snapshot(&[]);
|
||||
|
||||
assert_eq!(*service.my_addr_cache.lock().unwrap(), None);
|
||||
assert_eq!(*runtime.lease.lock().unwrap(), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reconcile_runtime_updates_public_lease_when_auto_enabled() {
|
||||
let public_addr = "2001:db8::123/64".parse().unwrap();
|
||||
let runtime = Arc::new(TestRuntime::new(true));
|
||||
|
||||
let service = Arc::new(PublicIpv6Service::new(
|
||||
runtime.clone(),
|
||||
std::sync::Weak::<PeerRpcManager>::new(),
|
||||
Arc::new(TestRouteControl {
|
||||
my_peer_id: 1,
|
||||
peers: Mutex::new(vec![PublicIpv6PeerRouteInfo {
|
||||
peer_id: 1,
|
||||
inst_id: Some(uuid::Uuid::from_u128(1)),
|
||||
is_provider: false,
|
||||
prefix: None,
|
||||
lease: Some(public_addr),
|
||||
reachable: true,
|
||||
}]),
|
||||
}),
|
||||
Arc::new(TestSyncTrigger),
|
||||
));
|
||||
|
||||
service.reconcile_runtime();
|
||||
|
||||
assert_eq!(*runtime.lease.lock().unwrap(), Some(public_addr));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_uses_explicit_host_capability() {
|
||||
let unsupported = PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: None,
|
||||
provider_supported: false,
|
||||
};
|
||||
assert_eq!(
|
||||
unsupported.validate(),
|
||||
Err(PublicIpv6ProviderConfigError::UnsupportedProvider)
|
||||
);
|
||||
|
||||
let disabled = PublicIpv6ProviderConfig {
|
||||
provider_enabled: false,
|
||||
..unsupported
|
||||
};
|
||||
assert!(disabled.validate().is_ok());
|
||||
assert!(!disabled.should_run_reconcile());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_rejects_non_global_prefixes() {
|
||||
for prefix in ["::1/128", "fe80::/64", "fd00::/48", "ff00::/8", "::/0"] {
|
||||
let config = PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: Some(prefix.parse().unwrap()),
|
||||
provider_supported: true,
|
||||
};
|
||||
assert!(matches!(
|
||||
config.validate(),
|
||||
Err(PublicIpv6ProviderConfigError::InvalidPrefix(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_config_accepts_global_prefix() {
|
||||
let config = PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: Some("2001:db8::/48".parse().unwrap()),
|
||||
provider_supported: true,
|
||||
};
|
||||
assert!(config.validate().is_ok());
|
||||
assert!(config.should_run_reconcile());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_resolution_prefers_configured_prefix_without_detection() {
|
||||
let prefix = "2001:db8::/48".parse().unwrap();
|
||||
let config = PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: Some(prefix),
|
||||
provider_supported: true,
|
||||
};
|
||||
assert_eq!(
|
||||
resolve_public_ipv6_provider(config, Err("detection must be ignored".to_owned())),
|
||||
PublicIpv6ProviderResolution::Active(prefix)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_resolution_normalizes_auto_detection_results() {
|
||||
let config = PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: None,
|
||||
provider_supported: true,
|
||||
};
|
||||
let prefix = "2001:db8:1::/56".parse().unwrap();
|
||||
assert_eq!(
|
||||
resolve_public_ipv6_provider(config, Ok(Some(prefix))),
|
||||
PublicIpv6ProviderResolution::Active(prefix)
|
||||
);
|
||||
assert!(matches!(
|
||||
resolve_public_ipv6_provider(config, Ok(None)),
|
||||
PublicIpv6ProviderResolution::Pending(message)
|
||||
if message.contains("ipv6-public-addr-prefix")
|
||||
));
|
||||
assert_eq!(
|
||||
resolve_public_ipv6_provider(config, Err("host detection failed".to_owned())),
|
||||
PublicIpv6ProviderResolution::Pending("host detection failed".to_owned())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_resolution_rejects_invalid_configured_and_detected_prefixes() {
|
||||
let configured = PublicIpv6ProviderConfig {
|
||||
provider_enabled: true,
|
||||
configured_prefix: Some("fd00::/48".parse().unwrap()),
|
||||
provider_supported: true,
|
||||
};
|
||||
assert!(matches!(
|
||||
resolve_public_ipv6_provider(configured, Ok(None)),
|
||||
PublicIpv6ProviderResolution::Pending(message)
|
||||
if message.contains("configured prefix")
|
||||
));
|
||||
|
||||
let detected = PublicIpv6ProviderConfig {
|
||||
configured_prefix: None,
|
||||
..configured
|
||||
};
|
||||
assert!(matches!(
|
||||
resolve_public_ipv6_provider(
|
||||
detected,
|
||||
Ok(Some("fe80::/64".parse().unwrap()))
|
||||
),
|
||||
PublicIpv6ProviderResolution::Pending(message)
|
||||
if message.contains("detected prefix")
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,601 @@
|
||||
use std::{
|
||||
sync::{
|
||||
Arc, Weak,
|
||||
atomic::{AtomicBool, Ordering},
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use cidr::Ipv6Cidr;
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
config::peers::PublicIpv6ProviderConfig,
|
||||
config::runtime::CoreRuntimeConfigStore,
|
||||
peers::public_ipv6::{
|
||||
CorePublicIpv6Runtime, PublicIpv6ProviderResolution, resolve_public_ipv6_provider,
|
||||
},
|
||||
};
|
||||
|
||||
const DEFAULT_RECONCILE_INTERVAL: Duration = Duration::from_secs(5);
|
||||
const MAX_CONFIG_RETRIES: usize = 3;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct PublicIpv6NdpTarget {
|
||||
pub wan_interface: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub struct PublicIpv6PlatformObservation {
|
||||
pub detected_prefix: Option<Ipv6Cidr>,
|
||||
pub ndp_target: Option<PublicIpv6NdpTarget>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct PublicIpv6NdpDesired {
|
||||
pub prefix: Ipv6Cidr,
|
||||
pub target: PublicIpv6NdpTarget,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
enum PublicIpv6ProviderState {
|
||||
Disabled,
|
||||
Pending(String),
|
||||
Active {
|
||||
prefix: Ipv6Cidr,
|
||||
ndp_target: Option<PublicIpv6NdpTarget>,
|
||||
},
|
||||
}
|
||||
|
||||
impl PublicIpv6ProviderState {
|
||||
fn from_resolution(
|
||||
resolution: PublicIpv6ProviderResolution,
|
||||
ndp_target: Option<PublicIpv6NdpTarget>,
|
||||
) -> Self {
|
||||
match resolution {
|
||||
PublicIpv6ProviderResolution::Disabled => Self::Disabled,
|
||||
PublicIpv6ProviderResolution::Pending(error) => Self::Pending(error),
|
||||
PublicIpv6ProviderResolution::Active(prefix) => Self::Active { prefix, ndp_target },
|
||||
}
|
||||
}
|
||||
|
||||
fn advertised_prefix(&self) -> Option<Ipv6Cidr> {
|
||||
match self {
|
||||
Self::Active { prefix, .. } => Some(*prefix),
|
||||
Self::Disabled | Self::Pending(_) => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn ndp_desired(&self) -> Option<PublicIpv6NdpDesired> {
|
||||
match self {
|
||||
Self::Active {
|
||||
prefix,
|
||||
ndp_target: Some(target),
|
||||
} => Some(PublicIpv6NdpDesired {
|
||||
prefix: *prefix,
|
||||
target: target.clone(),
|
||||
}),
|
||||
Self::Disabled | Self::Pending(_) | Self::Active { .. } => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, thiserror::Error, PartialEq, Eq)]
|
||||
pub enum PublicIpv6PlatformError {
|
||||
#[error("public IPv6 platform adapter is unavailable")]
|
||||
Unavailable,
|
||||
#[error("{0}")]
|
||||
Failed(String),
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait PublicIpv6ProviderPlatform: Send + Sync + 'static {
|
||||
fn inspect(
|
||||
&self,
|
||||
config: PublicIpv6ProviderConfig,
|
||||
) -> Result<PublicIpv6PlatformObservation, PublicIpv6PlatformError>;
|
||||
|
||||
fn sync_ndp(
|
||||
&self,
|
||||
desired: Option<PublicIpv6NdpDesired>,
|
||||
) -> Result<(), PublicIpv6PlatformError>;
|
||||
|
||||
/// Waits for a platform state change that requires an immediate retry.
|
||||
/// Returns `false` when the event source has closed.
|
||||
async fn wait_for_change(&self) -> bool;
|
||||
}
|
||||
|
||||
struct PublicIpv6ProviderTask {
|
||||
cancel: CancellationToken,
|
||||
handle: tokio::task::JoinHandle<()>,
|
||||
}
|
||||
|
||||
pub struct PublicIpv6ProviderService {
|
||||
platform: Arc<dyn PublicIpv6ProviderPlatform>,
|
||||
runtime_config: CoreRuntimeConfigStore,
|
||||
runtime: Arc<CorePublicIpv6Runtime>,
|
||||
reconcile_interval: Duration,
|
||||
reconcile: Mutex<()>,
|
||||
last_state: std::sync::Mutex<Option<PublicIpv6ProviderState>>,
|
||||
task: Mutex<Option<PublicIpv6ProviderTask>>,
|
||||
closing: AtomicBool,
|
||||
}
|
||||
|
||||
#[cfg(feature = "public-ipv6-provider")]
|
||||
pub(crate) struct PublicIpv6ProviderRuntime {
|
||||
service: Option<Arc<PublicIpv6ProviderService>>,
|
||||
runtime_config: CoreRuntimeConfigStore,
|
||||
}
|
||||
|
||||
#[cfg(feature = "public-ipv6-provider")]
|
||||
impl PublicIpv6ProviderRuntime {
|
||||
pub(crate) fn new(
|
||||
platform: Option<Arc<dyn PublicIpv6ProviderPlatform>>,
|
||||
runtime_config: CoreRuntimeConfigStore,
|
||||
runtime: Arc<CorePublicIpv6Runtime>,
|
||||
) -> Self {
|
||||
let service = platform.map(|platform| {
|
||||
PublicIpv6ProviderService::new(platform, runtime_config.clone(), runtime)
|
||||
});
|
||||
Self {
|
||||
service,
|
||||
runtime_config,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn validate_before_start(&self) -> anyhow::Result<()> {
|
||||
let config = self.runtime_config.snapshot().services.public_ipv6_provider;
|
||||
config.validate().map_err(anyhow::Error::new)?;
|
||||
if config.provider_enabled && self.service.is_none() {
|
||||
anyhow::bail!("public IPv6 provider is enabled but no host adapter was provided");
|
||||
}
|
||||
if let Some(service) = &self.service {
|
||||
service.apply_config().await;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) async fn start(&self) {
|
||||
if let Some(service) = &self.service {
|
||||
service.start().await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn stop(&self) {
|
||||
if let Some(service) = &self.service {
|
||||
service.stop().await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn reconcile(&self) -> bool {
|
||||
let Some(service) = &self.service else {
|
||||
return false;
|
||||
};
|
||||
let applied = service.apply_config().await;
|
||||
service.start().await;
|
||||
applied
|
||||
}
|
||||
}
|
||||
|
||||
impl PublicIpv6ProviderService {
|
||||
pub fn new(
|
||||
platform: Arc<dyn PublicIpv6ProviderPlatform>,
|
||||
runtime_config: CoreRuntimeConfigStore,
|
||||
runtime: Arc<CorePublicIpv6Runtime>,
|
||||
) -> Arc<Self> {
|
||||
Self::new_with_interval(
|
||||
platform,
|
||||
runtime_config,
|
||||
runtime,
|
||||
DEFAULT_RECONCILE_INTERVAL,
|
||||
)
|
||||
}
|
||||
|
||||
fn new_with_interval(
|
||||
platform: Arc<dyn PublicIpv6ProviderPlatform>,
|
||||
runtime_config: CoreRuntimeConfigStore,
|
||||
runtime: Arc<CorePublicIpv6Runtime>,
|
||||
reconcile_interval: Duration,
|
||||
) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
platform,
|
||||
runtime_config,
|
||||
runtime,
|
||||
reconcile_interval,
|
||||
reconcile: Mutex::new(()),
|
||||
last_state: std::sync::Mutex::new(None),
|
||||
task: Mutex::new(None),
|
||||
closing: AtomicBool::new(false),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn reconcile_now(&self) -> bool {
|
||||
let _reconcile = self.reconcile.lock().await;
|
||||
if self.closing.load(Ordering::Acquire) {
|
||||
return false;
|
||||
}
|
||||
for attempt in 0..MAX_CONFIG_RETRIES {
|
||||
let config = self.runtime_config.snapshot().services.public_ipv6_provider;
|
||||
let observation = if config.provider_enabled && config.provider_supported {
|
||||
match self.platform.inspect(config) {
|
||||
Ok(observation) => Ok(observation),
|
||||
Err(PublicIpv6PlatformError::Unavailable) => return false,
|
||||
Err(PublicIpv6PlatformError::Failed(error)) => Err(error),
|
||||
}
|
||||
} else {
|
||||
Ok(PublicIpv6PlatformObservation::default())
|
||||
};
|
||||
let (detected_prefix, ndp_target) = match observation {
|
||||
Ok(observation) => (Ok(observation.detected_prefix), observation.ndp_target),
|
||||
Err(error) => (Err(error), None),
|
||||
};
|
||||
let next_state = PublicIpv6ProviderState::from_resolution(
|
||||
resolve_public_ipv6_provider(config, detected_prefix),
|
||||
ndp_target,
|
||||
);
|
||||
|
||||
if self.runtime_config.snapshot().services.public_ipv6_provider != config {
|
||||
tracing::debug!(
|
||||
attempt = attempt + 1,
|
||||
max_retries = MAX_CONFIG_RETRIES,
|
||||
"public IPv6 provider config changed during reconcile, retrying"
|
||||
);
|
||||
continue;
|
||||
}
|
||||
|
||||
let changed = self
|
||||
.runtime
|
||||
.set_provider_prefix(next_state.advertised_prefix());
|
||||
if let Err(error) = self.platform.sync_ndp(next_state.ndp_desired()) {
|
||||
match error {
|
||||
PublicIpv6PlatformError::Unavailable => return false,
|
||||
PublicIpv6PlatformError::Failed(error) => {
|
||||
tracing::warn!(%error, "failed to synchronize public IPv6 NDP state");
|
||||
}
|
||||
}
|
||||
}
|
||||
self.log_state_change(&next_state, changed);
|
||||
*self.last_state.lock().unwrap() = Some(next_state);
|
||||
return true;
|
||||
}
|
||||
|
||||
tracing::warn!(
|
||||
max_retries = MAX_CONFIG_RETRIES,
|
||||
"skipping public IPv6 provider reconcile because config kept changing"
|
||||
);
|
||||
true
|
||||
}
|
||||
|
||||
pub async fn apply_config(&self) -> bool {
|
||||
self.reconcile_now().await
|
||||
}
|
||||
|
||||
pub async fn start(self: &Arc<Self>) {
|
||||
let mut task = self.task.lock().await;
|
||||
let config = self.runtime_config.snapshot().services.public_ipv6_provider;
|
||||
if self.closing.load(Ordering::Acquire) || task.is_some() || !config.should_run_reconcile()
|
||||
{
|
||||
return;
|
||||
}
|
||||
|
||||
let cancel = CancellationToken::new();
|
||||
let task_cancel = cancel.clone();
|
||||
let service = Arc::downgrade(self);
|
||||
let platform = self.platform.clone();
|
||||
let handle = tokio::spawn(async move {
|
||||
Self::run(service, platform, task_cancel).await;
|
||||
});
|
||||
task.replace(PublicIpv6ProviderTask { cancel, handle });
|
||||
}
|
||||
|
||||
pub async fn stop(&self) {
|
||||
self.closing.store(true, Ordering::Release);
|
||||
let task = self.task.lock().await.take();
|
||||
if let Some(task) = task {
|
||||
task.cancel.cancel();
|
||||
if let Err(error) = task.handle.await {
|
||||
tracing::warn!(?error, "public IPv6 provider task failed during shutdown");
|
||||
}
|
||||
}
|
||||
let _reconcile = self.reconcile.lock().await;
|
||||
if let Err(error) = self.platform.sync_ndp(None)
|
||||
&& !matches!(error, PublicIpv6PlatformError::Unavailable)
|
||||
{
|
||||
tracing::warn!(%error, "failed to clean public IPv6 NDP state during shutdown");
|
||||
}
|
||||
}
|
||||
|
||||
fn log_state_change(&self, next_state: &PublicIpv6ProviderState, changed: bool) {
|
||||
let last_state = self.last_state.lock().unwrap();
|
||||
if last_state.as_ref() != Some(next_state) {
|
||||
match next_state {
|
||||
PublicIpv6ProviderState::Disabled if last_state.is_some() => {
|
||||
tracing::info!("public IPv6 provider disabled");
|
||||
}
|
||||
PublicIpv6ProviderState::Disabled => {}
|
||||
PublicIpv6ProviderState::Pending(reason) => {
|
||||
tracing::warn!(%reason, "public IPv6 provider not ready");
|
||||
}
|
||||
PublicIpv6ProviderState::Active { prefix, ndp_target } => {
|
||||
if let Some(target) = ndp_target {
|
||||
tracing::info!(
|
||||
%prefix,
|
||||
wan_interface = %target.wan_interface,
|
||||
"public IPv6 provider is active with NDP proxy"
|
||||
);
|
||||
} else {
|
||||
tracing::info!(%prefix, "public IPv6 provider is active");
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if changed {
|
||||
tracing::info!("public IPv6 provider runtime state changed");
|
||||
}
|
||||
}
|
||||
|
||||
async fn run(
|
||||
service: Weak<Self>,
|
||||
platform: Arc<dyn PublicIpv6ProviderPlatform>,
|
||||
cancel: CancellationToken,
|
||||
) {
|
||||
loop {
|
||||
let Some(service) = service.upgrade() else {
|
||||
let _ = platform.sync_ndp(None);
|
||||
return;
|
||||
};
|
||||
if !service.reconcile_now().await {
|
||||
let _ = service.platform.sync_ndp(None);
|
||||
return;
|
||||
}
|
||||
|
||||
let interval = service.reconcile_interval;
|
||||
drop(service);
|
||||
let should_continue = tokio::select! {
|
||||
_ = cancel.cancelled() => false,
|
||||
_ = crate::foundation::time::sleep(interval) => true,
|
||||
changed = platform.wait_for_change() => changed,
|
||||
};
|
||||
if !should_continue {
|
||||
let _ = platform.sync_ndp(None);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::{
|
||||
Mutex as StdMutex,
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
};
|
||||
|
||||
use tokio::sync::Notify;
|
||||
|
||||
use super::*;
|
||||
use crate::{
|
||||
config::peers::PeerRuntimeSnapshot,
|
||||
config::runtime::{CoreRuntimeConfig, CoreRuntimeConfigStore},
|
||||
peers::context::PeerPublicIpv6State,
|
||||
};
|
||||
|
||||
struct RecordingHost {
|
||||
observation: StdMutex<Result<PublicIpv6PlatformObservation, PublicIpv6PlatformError>>,
|
||||
inspect_calls: AtomicUsize,
|
||||
ndp_desired: StdMutex<Vec<Option<PublicIpv6NdpDesired>>>,
|
||||
change: Notify,
|
||||
}
|
||||
|
||||
impl Default for RecordingHost {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
observation: StdMutex::new(Ok(PublicIpv6PlatformObservation::default())),
|
||||
inspect_calls: AtomicUsize::new(0),
|
||||
ndp_desired: StdMutex::new(Vec::new()),
|
||||
change: Notify::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl PublicIpv6ProviderPlatform for RecordingHost {
|
||||
fn inspect(
|
||||
&self,
|
||||
_config: PublicIpv6ProviderConfig,
|
||||
) -> Result<PublicIpv6PlatformObservation, PublicIpv6PlatformError> {
|
||||
self.inspect_calls.fetch_add(1, Ordering::AcqRel);
|
||||
self.observation.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
async fn wait_for_change(&self) -> bool {
|
||||
self.change.notified().await;
|
||||
true
|
||||
}
|
||||
|
||||
fn sync_ndp(
|
||||
&self,
|
||||
desired: Option<PublicIpv6NdpDesired>,
|
||||
) -> Result<(), PublicIpv6PlatformError> {
|
||||
self.ndp_desired.lock().unwrap().push(desired);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_config(enabled: bool, prefix: Option<Ipv6Cidr>) -> PublicIpv6ProviderConfig {
|
||||
PublicIpv6ProviderConfig {
|
||||
provider_enabled: enabled,
|
||||
configured_prefix: prefix,
|
||||
provider_supported: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn runtime_config(config: PublicIpv6ProviderConfig) -> CoreRuntimeConfigStore {
|
||||
let services = CoreRuntimeConfig {
|
||||
public_ipv6_provider: config,
|
||||
..Default::default()
|
||||
};
|
||||
CoreRuntimeConfigStore::new(services, Arc::new(PeerRuntimeSnapshot::default()))
|
||||
}
|
||||
|
||||
fn runtime(
|
||||
config: PublicIpv6ProviderConfig,
|
||||
) -> (CoreRuntimeConfigStore, Arc<CorePublicIpv6Runtime>) {
|
||||
let config = runtime_config(config);
|
||||
let runtime = CorePublicIpv6Runtime::new(config.clone(), Arc::new(()), Arc::new(()));
|
||||
(config, runtime)
|
||||
}
|
||||
|
||||
async fn wait_for_calls(counter: &AtomicUsize, expected: usize) {
|
||||
crate::foundation::time::timeout(Duration::from_secs(1), async {
|
||||
while counter.load(Ordering::Acquire) < expected {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn starts_only_when_enabled_and_reacts_to_host_changes() {
|
||||
let host = Arc::new(RecordingHost::default());
|
||||
let (runtime_config, runtime) = runtime(provider_config(false, None));
|
||||
let service = PublicIpv6ProviderService::new_with_interval(
|
||||
host.clone(),
|
||||
runtime_config.clone(),
|
||||
runtime,
|
||||
Duration::from_secs(60),
|
||||
);
|
||||
|
||||
service.start().await;
|
||||
assert_eq!(host.inspect_calls.load(Ordering::Acquire), 0);
|
||||
|
||||
runtime_config.update_services(|services| {
|
||||
services.public_ipv6_provider =
|
||||
provider_config(true, Some("2001:db8::/48".parse().unwrap()));
|
||||
});
|
||||
service.start().await;
|
||||
wait_for_calls(&host.inspect_calls, 1).await;
|
||||
host.change.notify_one();
|
||||
wait_for_calls(&host.inspect_calls, 2).await;
|
||||
|
||||
service.stop().await;
|
||||
assert_eq!(host.ndp_desired.lock().unwrap().last(), Some(&None));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn does_not_reconcile_or_restart_after_stop() {
|
||||
let host = Arc::new(RecordingHost::default());
|
||||
let (runtime_config, runtime) = runtime(provider_config(
|
||||
true,
|
||||
Some("2001:db8::/48".parse().unwrap()),
|
||||
));
|
||||
let service = PublicIpv6ProviderService::new_with_interval(
|
||||
host.clone(),
|
||||
runtime_config,
|
||||
runtime,
|
||||
Duration::from_secs(60),
|
||||
);
|
||||
|
||||
service.stop().await;
|
||||
assert!(!service.reconcile_now().await);
|
||||
service.start().await;
|
||||
|
||||
assert_eq!(host.inspect_calls.load(Ordering::Acquire), 0);
|
||||
assert_eq!(host.ndp_desired.lock().unwrap().as_slice(), &[None]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolves_observation_and_publishes_ndp_desired_state() {
|
||||
let host = Arc::new(RecordingHost::default());
|
||||
let prefix = "2001:db8::/48".parse().unwrap();
|
||||
let target = PublicIpv6NdpTarget {
|
||||
wan_interface: "wan0".to_owned(),
|
||||
};
|
||||
*host.observation.lock().unwrap() = Ok(PublicIpv6PlatformObservation {
|
||||
detected_prefix: Some(prefix),
|
||||
ndp_target: Some(target.clone()),
|
||||
});
|
||||
let (runtime_config, runtime) = runtime(provider_config(true, None));
|
||||
let service = PublicIpv6ProviderService::new(host.clone(), runtime_config, runtime.clone());
|
||||
|
||||
assert!(service.reconcile_now().await);
|
||||
|
||||
assert_eq!(runtime.advertised_ipv6_public_addr_prefix(), Some(prefix));
|
||||
assert!(runtime.public_ipv6_provider_enabled());
|
||||
assert_eq!(
|
||||
host.ndp_desired.lock().unwrap().as_slice(),
|
||||
&[Some(PublicIpv6NdpDesired { prefix, target })]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn turns_platform_failure_into_pending_provider_state() {
|
||||
let host = Arc::new(RecordingHost::default());
|
||||
*host.observation.lock().unwrap() = Err(PublicIpv6PlatformError::Failed(
|
||||
"route query failed".to_owned(),
|
||||
));
|
||||
let (runtime_config, runtime) = runtime(provider_config(true, None));
|
||||
let service = PublicIpv6ProviderService::new(host.clone(), runtime_config, runtime.clone());
|
||||
|
||||
assert!(service.reconcile_now().await);
|
||||
|
||||
assert_eq!(runtime.advertised_ipv6_public_addr_prefix(), None);
|
||||
assert!(!runtime.public_ipv6_provider_enabled());
|
||||
assert_eq!(host.ndp_desired.lock().unwrap().as_slice(), &[None]);
|
||||
}
|
||||
|
||||
struct ReconfiguringHost {
|
||||
runtime_config: CoreRuntimeConfigStore,
|
||||
replacement: PublicIpv6ProviderConfig,
|
||||
inspect_calls: AtomicUsize,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl PublicIpv6ProviderPlatform for ReconfiguringHost {
|
||||
fn inspect(
|
||||
&self,
|
||||
_config: PublicIpv6ProviderConfig,
|
||||
) -> Result<PublicIpv6PlatformObservation, PublicIpv6PlatformError> {
|
||||
if self.inspect_calls.fetch_add(1, Ordering::AcqRel) == 0 {
|
||||
self.runtime_config.update_services(|services| {
|
||||
services.public_ipv6_provider = self.replacement;
|
||||
});
|
||||
}
|
||||
Ok(PublicIpv6PlatformObservation::default())
|
||||
}
|
||||
|
||||
fn sync_ndp(
|
||||
&self,
|
||||
_desired: Option<PublicIpv6NdpDesired>,
|
||||
) -> Result<(), PublicIpv6PlatformError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn wait_for_change(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retries_when_config_changes_during_platform_inspection() {
|
||||
let first = provider_config(true, Some("2001:db8:1::/48".parse().unwrap()));
|
||||
let replacement = provider_config(true, Some("2001:db8:2::/48".parse().unwrap()));
|
||||
let (runtime_config, runtime) = runtime(first);
|
||||
let host = Arc::new(ReconfiguringHost {
|
||||
runtime_config: runtime_config.clone(),
|
||||
replacement,
|
||||
inspect_calls: AtomicUsize::new(0),
|
||||
});
|
||||
let service = PublicIpv6ProviderService::new(host.clone(), runtime_config, runtime.clone());
|
||||
|
||||
assert!(service.reconcile_now().await);
|
||||
|
||||
assert_eq!(host.inspect_calls.load(Ordering::Acquire), 2);
|
||||
assert_eq!(
|
||||
runtime.advertised_ipv6_public_addr_prefix(),
|
||||
replacement.configured_prefix
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,814 @@
|
||||
//! Lease-driven public IPv6 service: the lease allocator, the per-instance
|
||||
//! service driving acquisition/renewal, and the RPC server serving lease
|
||||
//! requests from client peers.
|
||||
|
||||
use std::{
|
||||
collections::{BTreeMap, BTreeSet, HashMap, HashSet},
|
||||
net::Ipv6Addr,
|
||||
sync::{Arc, Weak},
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
|
||||
use cidr::{Ipv6Cidr, Ipv6Inet};
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
peers::peer_rpc::PeerRpcManager,
|
||||
proto::{
|
||||
common::Void,
|
||||
peer_rpc::{
|
||||
AcquireIpv6PublicAddrLeaseRequest, GetIpv6PublicAddrLeaseRequest,
|
||||
Ipv6PublicAddrLeaseReply, PublicIpv6AddrRpc, PublicIpv6AddrRpcClientFactory,
|
||||
ReleaseIpv6PublicAddrLeaseRequest, RenewIpv6PublicAddrLeaseRequest,
|
||||
},
|
||||
rpc_types::{
|
||||
self,
|
||||
controller::{BaseController, Controller},
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
use super::{
|
||||
PublicIpv6PeerRouteInfo, PublicIpv6RouteControl, PublicIpv6Runtime, PublicIpv6SyncTrigger,
|
||||
};
|
||||
|
||||
// Use a longer lease with an early renew window to reduce steady-state RPC
|
||||
// churn while preserving enough margin for transient provider failures.
|
||||
static PUBLIC_IPV6_LEASE_TTL: Duration = Duration::from_secs(120);
|
||||
static PUBLIC_IPV6_RENEW_INTERVAL: Duration = Duration::from_secs(40);
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct PublicIpv6Provider {
|
||||
pub peer_id: PeerId,
|
||||
pub inst_id: uuid::Uuid,
|
||||
pub prefix: Ipv6Cidr,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub(crate) struct PublicIpv6ProviderLease {
|
||||
pub peer_id: PeerId,
|
||||
pub inst_id: uuid::Uuid,
|
||||
pub addr: Ipv6Inet,
|
||||
pub valid_until: SystemTime,
|
||||
pub reused: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct PublicIpv6ProviderState {
|
||||
provider: PublicIpv6Provider,
|
||||
leases: BTreeMap<uuid::Uuid, PublicIpv6ProviderLease>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct PublicIpv6ClientState {
|
||||
provider: PublicIpv6Provider,
|
||||
lease: PublicIpv6ProviderLease,
|
||||
last_error: Option<String>,
|
||||
}
|
||||
|
||||
pub(crate) struct PublicIpv6Service {
|
||||
runtime: Arc<dyn PublicIpv6Runtime>,
|
||||
peer_rpc: Weak<PeerRpcManager>,
|
||||
route_control: Arc<dyn PublicIpv6RouteControl>,
|
||||
sync_trigger: Arc<dyn PublicIpv6SyncTrigger>,
|
||||
|
||||
provider_state: std::sync::Mutex<Option<PublicIpv6ProviderState>>,
|
||||
client_state: std::sync::Mutex<Option<PublicIpv6ClientState>>,
|
||||
route_cache: std::sync::Mutex<BTreeSet<Ipv6Inet>>,
|
||||
pub(super) my_addr_cache: std::sync::Mutex<Option<Ipv6Inet>>,
|
||||
}
|
||||
|
||||
impl PublicIpv6Service {
|
||||
pub fn new(
|
||||
runtime: Arc<dyn PublicIpv6Runtime>,
|
||||
peer_rpc: Weak<PeerRpcManager>,
|
||||
route_control: Arc<dyn PublicIpv6RouteControl>,
|
||||
sync_trigger: Arc<dyn PublicIpv6SyncTrigger>,
|
||||
) -> Self {
|
||||
Self {
|
||||
runtime,
|
||||
peer_rpc,
|
||||
route_control,
|
||||
sync_trigger,
|
||||
provider_state: std::sync::Mutex::new(None),
|
||||
client_state: std::sync::Mutex::new(None),
|
||||
route_cache: std::sync::Mutex::new(BTreeSet::new()),
|
||||
my_addr_cache: std::sync::Mutex::new(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn rpc_server(self: &Arc<Self>) -> PublicIpv6AddrRpcServerImpl {
|
||||
PublicIpv6AddrRpcServerImpl {
|
||||
service: Arc::downgrade(self),
|
||||
}
|
||||
}
|
||||
|
||||
fn my_peer_id(&self) -> PeerId {
|
||||
self.route_control.my_peer_id()
|
||||
}
|
||||
|
||||
fn selected_provider(&self) -> Option<PublicIpv6Provider> {
|
||||
Self::selected_provider_from_snapshot(&self.route_control.peer_route_snapshot())
|
||||
}
|
||||
|
||||
fn current_provider_state(&self) -> Option<PublicIpv6ProviderState> {
|
||||
self.provider_state.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
fn current_client_state(&self) -> Option<PublicIpv6ClientState> {
|
||||
self.client_state.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
fn set_provider_state(&self, next: Option<PublicIpv6ProviderState>) -> bool {
|
||||
let mut guard = self.provider_state.lock().unwrap();
|
||||
if *guard == next {
|
||||
return false;
|
||||
}
|
||||
*guard = next;
|
||||
true
|
||||
}
|
||||
|
||||
fn set_client_state(&self, next: Option<PublicIpv6ClientState>) -> bool {
|
||||
let mut guard = self.client_state.lock().unwrap();
|
||||
if *guard == next {
|
||||
return false;
|
||||
}
|
||||
*guard = next;
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn selected_provider_from_snapshot(
|
||||
peers: &[PublicIpv6PeerRouteInfo],
|
||||
) -> Option<PublicIpv6Provider> {
|
||||
peers
|
||||
.iter()
|
||||
.filter(|info| info.is_provider)
|
||||
.filter(|info| info.reachable)
|
||||
.filter_map(|info| {
|
||||
Some(PublicIpv6Provider {
|
||||
peer_id: info.peer_id,
|
||||
inst_id: info.inst_id?,
|
||||
prefix: info.prefix?,
|
||||
})
|
||||
})
|
||||
.min_by_key(|provider| provider.inst_id)
|
||||
}
|
||||
|
||||
fn clear_provider_state_if_provider_changed(
|
||||
&self,
|
||||
provider: Option<&PublicIpv6Provider>,
|
||||
) -> bool {
|
||||
let current = self.current_provider_state();
|
||||
let should_clear = current
|
||||
.as_ref()
|
||||
.is_some_and(|state| provider != Some(&state.provider));
|
||||
should_clear && self.set_provider_state(None)
|
||||
}
|
||||
|
||||
fn clear_client_state_if_provider_changed(
|
||||
&self,
|
||||
provider: Option<&PublicIpv6Provider>,
|
||||
) -> bool {
|
||||
let current = self.current_client_state();
|
||||
let should_clear = current
|
||||
.as_ref()
|
||||
.is_some_and(|state| provider != Some(&state.provider));
|
||||
should_clear && self.set_client_state(None)
|
||||
}
|
||||
|
||||
fn collect_runtime_from_snapshot(
|
||||
&self,
|
||||
peers: &[PublicIpv6PeerRouteInfo],
|
||||
) -> (Option<Ipv6Inet>, BTreeSet<Ipv6Inet>) {
|
||||
let mut my_addr = self.current_client_state().map(|state| state.lease.addr);
|
||||
let mut routes = BTreeSet::new();
|
||||
|
||||
for info in peers {
|
||||
let Some(lease) = info.lease else {
|
||||
continue;
|
||||
};
|
||||
|
||||
if info.peer_id == self.my_peer_id() {
|
||||
my_addr = Some(lease);
|
||||
continue;
|
||||
}
|
||||
|
||||
if info.reachable {
|
||||
routes.insert(lease);
|
||||
}
|
||||
}
|
||||
|
||||
(my_addr, routes)
|
||||
}
|
||||
|
||||
pub(super) fn reconcile_runtime_from_snapshot(&self, peers: &[PublicIpv6PeerRouteInfo]) {
|
||||
let (mut my_addr, routes) = self.collect_runtime_from_snapshot(peers);
|
||||
if !self.runtime.ipv6_public_addr_auto() {
|
||||
my_addr = None;
|
||||
}
|
||||
|
||||
let mut cached_my_addr = self.my_addr_cache.lock().unwrap();
|
||||
if *cached_my_addr != my_addr {
|
||||
let old = *cached_my_addr;
|
||||
*cached_my_addr = my_addr;
|
||||
self.runtime.public_ipv6_lease_changed(old, my_addr);
|
||||
}
|
||||
drop(cached_my_addr);
|
||||
|
||||
let mut cached_routes = self.route_cache.lock().unwrap();
|
||||
if *cached_routes != routes {
|
||||
let added = routes
|
||||
.difference(&cached_routes)
|
||||
.copied()
|
||||
.collect::<Vec<_>>();
|
||||
let removed = cached_routes
|
||||
.difference(&routes)
|
||||
.copied()
|
||||
.collect::<Vec<_>>();
|
||||
*cached_routes = routes;
|
||||
self.runtime.public_ipv6_routes_changed(added, removed);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn reconcile_runtime(&self) {
|
||||
let peers = self.route_control.peer_route_snapshot();
|
||||
self.reconcile_runtime_from_snapshot(&peers);
|
||||
}
|
||||
|
||||
pub fn handle_route_change(&self) -> bool {
|
||||
let peers = self.route_control.peer_route_snapshot();
|
||||
let provider = Self::selected_provider_from_snapshot(&peers);
|
||||
let _provider_changed = self.clear_provider_state_if_provider_changed(provider.as_ref());
|
||||
let client_changed = self.clear_client_state_if_provider_changed(provider.as_ref());
|
||||
|
||||
let peer_info_changed = if client_changed {
|
||||
self.publish_current_client_lease()
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
// When client state changed, publish_current_client_lease() mutated the
|
||||
// local peer info synchronously, so the pre-update snapshot is stale for
|
||||
// this node's own entry. Re-fetch to avoid reconciling against old data.
|
||||
if client_changed {
|
||||
self.reconcile_runtime();
|
||||
} else {
|
||||
self.reconcile_runtime_from_snapshot(&peers);
|
||||
}
|
||||
peer_info_changed
|
||||
}
|
||||
|
||||
fn publish_current_client_lease(&self) -> bool {
|
||||
self.route_control.publish_self_public_ipv6_lease(
|
||||
self.current_client_state()
|
||||
.as_ref()
|
||||
.map(|state| state.lease.addr),
|
||||
)
|
||||
}
|
||||
|
||||
fn clear_client_lease_state(&self, mut state_changed: bool) -> bool {
|
||||
state_changed |= self.set_client_state(None);
|
||||
let peer_info_changed = if state_changed {
|
||||
self.route_control.publish_self_public_ipv6_lease(None)
|
||||
} else {
|
||||
false
|
||||
};
|
||||
if state_changed {
|
||||
// publish_self_public_ipv6_lease mutated the local peer info above,
|
||||
// so the snapshot passed in is stale for this node.
|
||||
self.reconcile_runtime();
|
||||
}
|
||||
peer_info_changed
|
||||
}
|
||||
|
||||
fn build_lease_reply(
|
||||
provider: &PublicIpv6Provider,
|
||||
lease: Option<&PublicIpv6ProviderLease>,
|
||||
error_msg: Option<String>,
|
||||
) -> Ipv6PublicAddrLeaseReply {
|
||||
Ipv6PublicAddrLeaseReply {
|
||||
provider_peer_id: provider.peer_id,
|
||||
provider_inst_id: Some(provider.inst_id.into()),
|
||||
provider_prefix: Some(
|
||||
Ipv6Inet::new(
|
||||
provider.prefix.first_address(),
|
||||
provider.prefix.network_length(),
|
||||
)
|
||||
.unwrap()
|
||||
.into(),
|
||||
),
|
||||
leased_addr: lease.map(|lease| lease.addr.into()),
|
||||
valid_until: lease.map(|lease| lease.valid_until.into()),
|
||||
reused: lease.map(|lease| lease.reused).unwrap_or(false),
|
||||
error_msg,
|
||||
}
|
||||
}
|
||||
|
||||
async fn collect_reserved_addrs(&self, prefix: Ipv6Cidr) -> HashSet<Ipv6Addr> {
|
||||
self.runtime
|
||||
.collect_reserved_public_ipv6_addrs(prefix)
|
||||
.await
|
||||
}
|
||||
|
||||
fn prune_expired_leases(
|
||||
provider: &PublicIpv6Provider,
|
||||
current: Option<PublicIpv6ProviderState>,
|
||||
) -> PublicIpv6ProviderState {
|
||||
let mut state = current.unwrap_or_else(|| PublicIpv6ProviderState {
|
||||
provider: provider.clone(),
|
||||
leases: BTreeMap::new(),
|
||||
});
|
||||
state.provider = provider.clone();
|
||||
let now = SystemTime::now();
|
||||
state.leases.retain(|_, lease| lease.valid_until > now);
|
||||
state
|
||||
}
|
||||
|
||||
async fn acquire_lease(
|
||||
&self,
|
||||
requester_peer_id: PeerId,
|
||||
requester_inst_id: uuid::Uuid,
|
||||
renew_only: bool,
|
||||
requested_addr: Option<Ipv6Inet>,
|
||||
) -> Result<PublicIpv6ProviderLease, String> {
|
||||
let provider = self
|
||||
.selected_provider()
|
||||
.ok_or_else(|| "no active ipv6 public address provider".to_string())?;
|
||||
if provider.peer_id != self.my_peer_id() {
|
||||
return Err("this peer is not the selected ipv6 public address provider".to_string());
|
||||
}
|
||||
|
||||
let mut state = Self::prune_expired_leases(&provider, self.current_provider_state());
|
||||
if let Some(existing) = state.leases.get_mut(&requester_inst_id) {
|
||||
if requested_addr.is_some() && requested_addr != Some(existing.addr) {
|
||||
return Err("requested lease does not match the active allocation".to_string());
|
||||
}
|
||||
existing.peer_id = requester_peer_id;
|
||||
existing.valid_until = SystemTime::now() + PUBLIC_IPV6_LEASE_TTL;
|
||||
existing.reused = true;
|
||||
let lease = existing.clone();
|
||||
self.set_provider_state(Some(state));
|
||||
return Ok(lease);
|
||||
}
|
||||
|
||||
if renew_only {
|
||||
return Err("lease not found".to_string());
|
||||
}
|
||||
|
||||
let mut reserved = self.collect_reserved_addrs(provider.prefix).await;
|
||||
let old_map = state
|
||||
.leases
|
||||
.iter()
|
||||
.map(|(inst_id, lease)| {
|
||||
reserved.insert(lease.addr.address());
|
||||
(*inst_id, lease.addr)
|
||||
})
|
||||
.collect::<HashMap<_, _>>();
|
||||
|
||||
let mut allocated =
|
||||
allocate_public_ipv6_leases(provider.prefix, &[requester_inst_id], &reserved, &old_map);
|
||||
let Some(mut lease) = allocated.pop() else {
|
||||
return Err(format!(
|
||||
"no free ipv6 address left in provider prefix {}",
|
||||
provider.prefix
|
||||
));
|
||||
};
|
||||
lease.peer_id = requester_peer_id;
|
||||
lease.valid_until = SystemTime::now() + PUBLIC_IPV6_LEASE_TTL;
|
||||
|
||||
state.leases.insert(requester_inst_id, lease.clone());
|
||||
self.set_provider_state(Some(state));
|
||||
Ok(lease)
|
||||
}
|
||||
|
||||
fn release_lease(&self, requester_peer_id: PeerId, requester_inst_id: uuid::Uuid) -> bool {
|
||||
let Some(provider) = self.selected_provider() else {
|
||||
return false;
|
||||
};
|
||||
if provider.peer_id != self.my_peer_id() {
|
||||
return false;
|
||||
}
|
||||
|
||||
let mut state = Self::prune_expired_leases(&provider, self.current_provider_state());
|
||||
let removed = state
|
||||
.leases
|
||||
.get(&requester_inst_id)
|
||||
.map(|lease| lease.peer_id == requester_peer_id)
|
||||
.unwrap_or(false);
|
||||
if !removed {
|
||||
return false;
|
||||
}
|
||||
|
||||
state.leases.remove(&requester_inst_id);
|
||||
self.set_provider_state(Some(state))
|
||||
}
|
||||
|
||||
fn get_lease(
|
||||
&self,
|
||||
requester_peer_id: PeerId,
|
||||
requester_inst_id: uuid::Uuid,
|
||||
requested_addr: Option<Ipv6Inet>,
|
||||
) -> Result<(PublicIpv6Provider, PublicIpv6ProviderLease), String> {
|
||||
let provider = self
|
||||
.selected_provider()
|
||||
.ok_or_else(|| "no active ipv6 public address provider".to_string())?;
|
||||
if provider.peer_id != self.my_peer_id() {
|
||||
return Err("this peer is not the selected ipv6 public address provider".to_string());
|
||||
}
|
||||
|
||||
let state = Self::prune_expired_leases(&provider, self.current_provider_state());
|
||||
let Some(lease) = state.leases.get(&requester_inst_id) else {
|
||||
return Err("lease not found".to_string());
|
||||
};
|
||||
if lease.peer_id != requester_peer_id {
|
||||
return Err("lease owner mismatch".to_string());
|
||||
}
|
||||
if requested_addr.is_some() && requested_addr != Some(lease.addr) {
|
||||
return Err("requested lease does not match the active allocation".to_string());
|
||||
}
|
||||
Ok((provider, lease.clone()))
|
||||
}
|
||||
|
||||
pub async fn gc_provider_leases(&self) {
|
||||
let peers = self.route_control.peer_route_snapshot();
|
||||
let provider = Self::selected_provider_from_snapshot(&peers);
|
||||
self.clear_provider_state_if_provider_changed(provider.as_ref());
|
||||
|
||||
let Some(provider) = provider else {
|
||||
return;
|
||||
};
|
||||
if provider.peer_id != self.my_peer_id() {
|
||||
return;
|
||||
}
|
||||
|
||||
let state = Self::prune_expired_leases(&provider, self.current_provider_state());
|
||||
self.set_provider_state(Some(state));
|
||||
}
|
||||
|
||||
pub async fn sync_client_state(&self) -> bool {
|
||||
if !self.runtime.ipv6_public_addr_auto() {
|
||||
return self
|
||||
.clear_client_lease_state(self.clear_client_state_if_provider_changed(None));
|
||||
}
|
||||
|
||||
let peers = self.route_control.peer_route_snapshot();
|
||||
let provider = Self::selected_provider_from_snapshot(&peers);
|
||||
self.clear_provider_state_if_provider_changed(provider.as_ref());
|
||||
let state_changed = self.clear_client_state_if_provider_changed(provider.as_ref());
|
||||
|
||||
let Some(provider) = provider else {
|
||||
return self.clear_client_lease_state(state_changed);
|
||||
};
|
||||
|
||||
if provider.peer_id == self.my_peer_id() {
|
||||
return self.clear_client_lease_state(state_changed);
|
||||
}
|
||||
|
||||
let current = self.current_client_state();
|
||||
let need_rpc = current.as_ref().is_none_or(|state| {
|
||||
state.provider != provider
|
||||
|| state.lease.valid_until <= SystemTime::now() + PUBLIC_IPV6_RENEW_INTERVAL
|
||||
});
|
||||
|
||||
if !need_rpc {
|
||||
if state_changed {
|
||||
self.reconcile_runtime();
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
let Some(peer_rpc) = self.peer_rpc.upgrade() else {
|
||||
if state_changed {
|
||||
self.reconcile_runtime();
|
||||
}
|
||||
return false;
|
||||
};
|
||||
|
||||
let mut ctrl = BaseController::default();
|
||||
ctrl.set_timeout_ms(3000);
|
||||
let rpc_stub = peer_rpc
|
||||
.rpc_client()
|
||||
.scoped_client::<PublicIpv6AddrRpcClientFactory<BaseController>>(
|
||||
self.my_peer_id(),
|
||||
provider.peer_id,
|
||||
self.runtime.network_name(),
|
||||
);
|
||||
|
||||
let inst_id = self.runtime.instance_id();
|
||||
let reply = if let Some(state) = current.as_ref().filter(|state| state.provider == provider)
|
||||
{
|
||||
match rpc_stub
|
||||
.renew_lease(
|
||||
ctrl.clone(),
|
||||
RenewIpv6PublicAddrLeaseRequest {
|
||||
peer_id: self.my_peer_id(),
|
||||
inst_id: Some(inst_id.into()),
|
||||
leased_addr: Some(state.lease.addr.into()),
|
||||
},
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(reply) if reply.error_msg.is_none() => Ok(reply),
|
||||
Ok(_) | Err(_) => {
|
||||
rpc_stub
|
||||
.acquire_lease(
|
||||
ctrl.clone(),
|
||||
AcquireIpv6PublicAddrLeaseRequest {
|
||||
peer_id: self.my_peer_id(),
|
||||
inst_id: Some(inst_id.into()),
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
} else {
|
||||
rpc_stub
|
||||
.acquire_lease(
|
||||
ctrl,
|
||||
AcquireIpv6PublicAddrLeaseRequest {
|
||||
peer_id: self.my_peer_id(),
|
||||
inst_id: Some(inst_id.into()),
|
||||
},
|
||||
)
|
||||
.await
|
||||
};
|
||||
|
||||
let mut state_changed = state_changed;
|
||||
|
||||
match reply {
|
||||
Ok(reply) if reply.error_msg.is_none() => {
|
||||
let Some(leased_addr) = reply.leased_addr.map(Into::into) else {
|
||||
return false;
|
||||
};
|
||||
let valid_until = reply
|
||||
.valid_until
|
||||
.and_then(|ts| SystemTime::try_from(ts).ok())
|
||||
.unwrap_or_else(|| SystemTime::now() + PUBLIC_IPV6_LEASE_TTL);
|
||||
let next_state = PublicIpv6ClientState {
|
||||
provider: provider.clone(),
|
||||
lease: PublicIpv6ProviderLease {
|
||||
peer_id: self.my_peer_id(),
|
||||
inst_id,
|
||||
addr: leased_addr,
|
||||
valid_until,
|
||||
reused: reply.reused,
|
||||
},
|
||||
last_error: None,
|
||||
};
|
||||
state_changed |= self.set_client_state(Some(next_state));
|
||||
}
|
||||
Ok(_) | Err(_) => {
|
||||
let should_clear = current
|
||||
.as_ref()
|
||||
.map(|state| state.lease.valid_until <= SystemTime::now())
|
||||
.unwrap_or(true);
|
||||
if should_clear {
|
||||
state_changed |= self.set_client_state(None);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let peer_info_changed = if state_changed {
|
||||
self.publish_current_client_lease()
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
if state_changed {
|
||||
self.reconcile_runtime();
|
||||
}
|
||||
|
||||
peer_info_changed
|
||||
}
|
||||
|
||||
pub async fn provider_gc_routine(self: Arc<Self>) {
|
||||
if !self.runtime.ipv6_public_addr_provider() {
|
||||
return;
|
||||
}
|
||||
loop {
|
||||
crate::foundation::time::sleep(Duration::from_secs(15)).await;
|
||||
self.gc_provider_leases().await;
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn client_routine(self: Arc<Self>) {
|
||||
loop {
|
||||
if self.sync_client_state().await {
|
||||
self.sync_trigger.sync_now("sync_public_ipv6_client_state");
|
||||
}
|
||||
crate::foundation::time::sleep(Duration::from_secs(5)).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn list_routes(&self) -> BTreeSet<Ipv6Inet> {
|
||||
self.route_cache.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
pub fn my_addr(&self) -> Option<Ipv6Inet> {
|
||||
*self.my_addr_cache.lock().unwrap()
|
||||
}
|
||||
|
||||
pub fn provider_peer_id_for_client(&self) -> Option<PeerId> {
|
||||
self.current_client_state()
|
||||
.map(|state| state.provider.peer_id)
|
||||
}
|
||||
|
||||
pub fn local_provider_state(
|
||||
&self,
|
||||
) -> Option<(PublicIpv6Provider, Vec<PublicIpv6ProviderLease>)> {
|
||||
let provider = self.selected_provider()?;
|
||||
if provider.peer_id != self.my_peer_id() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let state = Self::prune_expired_leases(&provider, self.current_provider_state());
|
||||
let mut leases = state.leases.into_values().collect::<Vec<_>>();
|
||||
leases.sort_by_key(|lease| (lease.peer_id, lease.inst_id, lease.addr));
|
||||
Some((provider, leases))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct PublicIpv6AddrRpcServerImpl {
|
||||
service: Weak<PublicIpv6Service>,
|
||||
}
|
||||
|
||||
impl PublicIpv6AddrRpcServerImpl {
|
||||
fn selected_provider(
|
||||
service: &PublicIpv6Service,
|
||||
) -> rpc_types::error::Result<PublicIpv6Provider> {
|
||||
service
|
||||
.selected_provider()
|
||||
.ok_or_else(|| anyhow::anyhow!("provider not available").into())
|
||||
}
|
||||
|
||||
fn build_error_reply(
|
||||
service: &PublicIpv6Service,
|
||||
error_msg: String,
|
||||
) -> rpc_types::error::Result<Ipv6PublicAddrLeaseReply> {
|
||||
Ok(PublicIpv6Service::build_lease_reply(
|
||||
&Self::selected_provider(service)?,
|
||||
None,
|
||||
Some(error_msg),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl PublicIpv6AddrRpc for PublicIpv6AddrRpcServerImpl {
|
||||
type Controller = BaseController;
|
||||
|
||||
async fn acquire_lease(
|
||||
&self,
|
||||
_: BaseController,
|
||||
request: AcquireIpv6PublicAddrLeaseRequest,
|
||||
) -> rpc_types::error::Result<Ipv6PublicAddrLeaseReply> {
|
||||
let Some(service) = self.service.upgrade() else {
|
||||
return Err(anyhow::anyhow!("public ipv6 service stopped").into());
|
||||
};
|
||||
let inst_id: uuid::Uuid = request
|
||||
.inst_id
|
||||
.ok_or_else(|| anyhow::anyhow!("inst_id is required"))?
|
||||
.into();
|
||||
|
||||
match service
|
||||
.acquire_lease(request.peer_id, inst_id, false, None)
|
||||
.await
|
||||
{
|
||||
Ok(lease) => Ok(PublicIpv6Service::build_lease_reply(
|
||||
&Self::selected_provider(&service)?,
|
||||
Some(&lease),
|
||||
None,
|
||||
)),
|
||||
Err(error_msg) => Self::build_error_reply(&service, error_msg),
|
||||
}
|
||||
}
|
||||
|
||||
async fn renew_lease(
|
||||
&self,
|
||||
_: BaseController,
|
||||
request: RenewIpv6PublicAddrLeaseRequest,
|
||||
) -> rpc_types::error::Result<Ipv6PublicAddrLeaseReply> {
|
||||
let Some(service) = self.service.upgrade() else {
|
||||
return Err(anyhow::anyhow!("public ipv6 service stopped").into());
|
||||
};
|
||||
let inst_id: uuid::Uuid = request
|
||||
.inst_id
|
||||
.ok_or_else(|| anyhow::anyhow!("inst_id is required"))?
|
||||
.into();
|
||||
let requested_addr = request.leased_addr.map(Into::into);
|
||||
|
||||
match service
|
||||
.acquire_lease(request.peer_id, inst_id, true, requested_addr)
|
||||
.await
|
||||
{
|
||||
Ok(lease) => Ok(PublicIpv6Service::build_lease_reply(
|
||||
&Self::selected_provider(&service)?,
|
||||
Some(&lease),
|
||||
None,
|
||||
)),
|
||||
Err(error_msg) => Self::build_error_reply(&service, error_msg),
|
||||
}
|
||||
}
|
||||
|
||||
async fn release_lease(
|
||||
&self,
|
||||
_: BaseController,
|
||||
request: ReleaseIpv6PublicAddrLeaseRequest,
|
||||
) -> rpc_types::error::Result<Void> {
|
||||
let Some(service) = self.service.upgrade() else {
|
||||
return Err(anyhow::anyhow!("public ipv6 service stopped").into());
|
||||
};
|
||||
let inst_id: uuid::Uuid = request
|
||||
.inst_id
|
||||
.ok_or_else(|| anyhow::anyhow!("inst_id is required"))?
|
||||
.into();
|
||||
service.release_lease(request.peer_id, inst_id);
|
||||
Ok(Default::default())
|
||||
}
|
||||
|
||||
async fn get_lease(
|
||||
&self,
|
||||
_: BaseController,
|
||||
request: GetIpv6PublicAddrLeaseRequest,
|
||||
) -> rpc_types::error::Result<Ipv6PublicAddrLeaseReply> {
|
||||
let Some(service) = self.service.upgrade() else {
|
||||
return Err(anyhow::anyhow!("public ipv6 service stopped").into());
|
||||
};
|
||||
let inst_id: uuid::Uuid = request
|
||||
.inst_id
|
||||
.ok_or_else(|| anyhow::anyhow!("inst_id is required"))?
|
||||
.into();
|
||||
match service.get_lease(request.peer_id, inst_id, None) {
|
||||
Ok((provider, lease)) => Ok(PublicIpv6Service::build_lease_reply(
|
||||
&provider,
|
||||
Some(&lease),
|
||||
None,
|
||||
)),
|
||||
Err(error_msg) => Self::build_error_reply(&service, error_msg),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn allocate_public_ipv6_leases(
|
||||
prefix: Ipv6Cidr,
|
||||
auto_peer_ids: &[uuid::Uuid],
|
||||
reserved: &HashSet<Ipv6Addr>,
|
||||
old_map: &HashMap<uuid::Uuid, Ipv6Inet>,
|
||||
) -> Vec<PublicIpv6ProviderLease> {
|
||||
let prefix_len = prefix.network_length();
|
||||
let host_bits = 128_u32.saturating_sub(prefix_len as u32);
|
||||
let max_offsets = if host_bits == 128 {
|
||||
None
|
||||
} else {
|
||||
Some(1_u128 << host_bits)
|
||||
};
|
||||
let network = u128::from(prefix.first_address());
|
||||
|
||||
let mut used_offsets = reserved
|
||||
.iter()
|
||||
.filter(|addr| prefix.contains(addr))
|
||||
.map(|addr| u128::from(*addr).saturating_sub(network))
|
||||
.collect::<HashSet<_>>();
|
||||
|
||||
let mut leases = Vec::with_capacity(auto_peer_ids.len());
|
||||
for inst_id in auto_peer_ids.iter().copied() {
|
||||
let addr = if let Some(existing) = old_map.get(&inst_id).copied()
|
||||
&& prefix.contains(&existing.address())
|
||||
&& used_offsets.insert(u128::from(existing.address()).saturating_sub(network))
|
||||
{
|
||||
existing
|
||||
} else {
|
||||
let Some(max_offsets) = max_offsets else {
|
||||
continue;
|
||||
};
|
||||
let usable_slots = max_offsets.saturating_sub(1);
|
||||
let offset = if usable_slots == 0 {
|
||||
used_offsets.insert(0).then_some(0)
|
||||
} else {
|
||||
let start_offset = (inst_id.as_u128() % usable_slots) + 1;
|
||||
(0..usable_slots)
|
||||
.map(|step| ((start_offset - 1 + step) % usable_slots) + 1)
|
||||
.find(|offset| used_offsets.insert(*offset))
|
||||
};
|
||||
let Some(offset) = offset else {
|
||||
break;
|
||||
};
|
||||
|
||||
Ipv6Inet::new(Ipv6Addr::from(network + offset), 128).unwrap()
|
||||
};
|
||||
|
||||
leases.push(PublicIpv6ProviderLease {
|
||||
peer_id: 0,
|
||||
inst_id,
|
||||
addr,
|
||||
valid_until: SystemTime::UNIX_EPOCH,
|
||||
reused: old_map
|
||||
.get(&inst_id)
|
||||
.map(|old| *old == addr)
|
||||
.unwrap_or(false),
|
||||
});
|
||||
}
|
||||
|
||||
leases
|
||||
}
|
||||
@@ -0,0 +1,753 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use dashmap::DashMap;
|
||||
use prost::Message;
|
||||
use quanta::Instant;
|
||||
use snow::params::NoiseParams;
|
||||
use tokio::sync::{Mutex, OwnedMutexGuard, oneshot};
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
foundation::time::{Duration, timeout},
|
||||
packet::{PacketType, ZCPacket},
|
||||
peers::{
|
||||
conn::{
|
||||
peer_map::PeerMap,
|
||||
peer_session::{PeerSession, PeerSessionAction, PeerSessionStore, SessionKey},
|
||||
},
|
||||
context::ArcPeerContext,
|
||||
error::Error,
|
||||
foreign_network::client::ForeignNetworkClient,
|
||||
route::NextHopPolicy,
|
||||
util::shrink_dashmap,
|
||||
},
|
||||
proto::peer_rpc::RoutePeerInfo,
|
||||
proto::peer_rpc::{PeerConnSessionActionPb, RelayNoiseMsg1Pb, RelayNoiseMsg2Pb},
|
||||
};
|
||||
|
||||
const RELAY_NOISE_VERSION: u32 = 1;
|
||||
const RELAY_NOISE_PROLOGUE: &[u8] = b"easytier-relay-noise";
|
||||
const HANDSHAKE_TIMEOUT_SECS: u64 = 5;
|
||||
const HANDSHAKE_RETRY_BASE_MS: u64 = 200;
|
||||
const HANDSHAKE_MAX_ATTEMPTS: u32 = 3;
|
||||
const MAX_PENDING_PACKETS_PER_PEER: usize = 32;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RelayPeerState {
|
||||
pub last_active_at: Instant,
|
||||
pub failure_count: u32,
|
||||
pub next_retry_at: Option<Instant>,
|
||||
}
|
||||
|
||||
impl Default for RelayPeerState {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
last_active_at: Instant::now(),
|
||||
failure_count: 0,
|
||||
next_retry_at: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RelayPeerMap {
|
||||
route_transport: Arc<dyn RelayRouteTransport>,
|
||||
context: ArcPeerContext,
|
||||
metric_network_name: String,
|
||||
my_peer_id: PeerId,
|
||||
peer_session_store: Arc<PeerSessionStore>,
|
||||
states: DashMap<PeerId, RelayPeerState>,
|
||||
pending_handshakes: DashMap<PeerId, oneshot::Sender<ZCPacket>>,
|
||||
handshake_locks: DashMap<PeerId, Arc<Mutex<()>>>,
|
||||
pub(crate) pending_packets: DashMap<PeerId, Vec<(ZCPacket, NextHopPolicy)>>,
|
||||
|
||||
is_secure_mode_enabled: bool,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(Arc)]
|
||||
pub trait RelayRouteTransport: Send + Sync {
|
||||
async fn get_route_peer_info(&self, peer_id: PeerId) -> Option<RoutePeerInfo>;
|
||||
|
||||
async fn send_msg_to_next_hop(
|
||||
&self,
|
||||
msg: ZCPacket,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<(), Error>;
|
||||
}
|
||||
|
||||
pub struct PeerMapRelayRouteTransport {
|
||||
peer_map: Arc<PeerMap>,
|
||||
foreign_network_client: Option<Arc<ForeignNetworkClient>>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl RelayRouteTransport for PeerMapRelayRouteTransport {
|
||||
async fn get_route_peer_info(&self, peer_id: PeerId) -> Option<RoutePeerInfo> {
|
||||
self.peer_map.get_route_peer_info(peer_id).await
|
||||
}
|
||||
|
||||
async fn send_msg_to_next_hop(
|
||||
&self,
|
||||
msg: ZCPacket,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<(), Error> {
|
||||
let Some(next_hop) = self.peer_map.get_gateway_peer_id(dst_peer_id, policy).await else {
|
||||
return Err(Error::RouteError(Some(format!(
|
||||
"next hop not found in route for peer {dst_peer_id:?}"
|
||||
))));
|
||||
};
|
||||
if self.peer_map.has_peer(next_hop) {
|
||||
self.peer_map.send_msg_directly(msg, next_hop).await
|
||||
} else if let Some(foreign_network_client) = &self.foreign_network_client {
|
||||
foreign_network_client.send_msg(msg, next_hop).await
|
||||
} else {
|
||||
Err(Error::RouteError(Some(format!(
|
||||
"next hop not found in direct peer map: {next_hop:?}"
|
||||
))))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_relay_peer_map(
|
||||
peer_map: Arc<PeerMap>,
|
||||
foreign_network_client: Option<Arc<ForeignNetworkClient>>,
|
||||
context: ArcPeerContext,
|
||||
my_peer_id: PeerId,
|
||||
peer_session_store: Arc<PeerSessionStore>,
|
||||
) -> Arc<RelayPeerMap> {
|
||||
RelayPeerMap::new(
|
||||
Arc::new(PeerMapRelayRouteTransport {
|
||||
peer_map,
|
||||
foreign_network_client,
|
||||
}),
|
||||
context,
|
||||
my_peer_id,
|
||||
peer_session_store,
|
||||
)
|
||||
}
|
||||
|
||||
impl RelayPeerMap {
|
||||
pub(crate) fn new(
|
||||
route_transport: Arc<dyn RelayRouteTransport>,
|
||||
context: ArcPeerContext,
|
||||
my_peer_id: PeerId,
|
||||
peer_session_store: Arc<PeerSessionStore>,
|
||||
) -> Arc<Self> {
|
||||
let is_secure_mode_enabled = context
|
||||
.secure_mode()
|
||||
.map(|cfg| cfg.enabled)
|
||||
.unwrap_or(false);
|
||||
let metric_network_name = context.network_name();
|
||||
Arc::new(Self {
|
||||
route_transport,
|
||||
context,
|
||||
metric_network_name,
|
||||
my_peer_id,
|
||||
peer_session_store,
|
||||
states: DashMap::new(),
|
||||
pending_handshakes: DashMap::new(),
|
||||
handshake_locks: DashMap::new(),
|
||||
pending_packets: DashMap::new(),
|
||||
is_secure_mode_enabled,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn is_secure_mode_enabled(&self) -> bool {
|
||||
self.is_secure_mode_enabled
|
||||
}
|
||||
|
||||
fn get_local_keypair(&self) -> Result<(Vec<u8>, Vec<u8>), Error> {
|
||||
let cfg = self
|
||||
.context
|
||||
.secure_mode()
|
||||
.ok_or_else(|| Error::RouteError(Some("secure mode config not set".to_string())))?;
|
||||
let private = cfg
|
||||
.private_key()
|
||||
.map_err(|e| Error::RouteError(Some(format!("invalid private key: {e:?}"))))?;
|
||||
let public = cfg
|
||||
.public_key()
|
||||
.map_err(|e| Error::RouteError(Some(format!("invalid public key: {e:?}"))))?;
|
||||
Ok((private.as_bytes().to_vec(), public.as_bytes().to_vec()))
|
||||
}
|
||||
|
||||
async fn get_remote_static_pubkey(&self, peer_id: PeerId) -> Result<Vec<u8>, Error> {
|
||||
let info = self
|
||||
.route_transport
|
||||
.get_route_peer_info(peer_id)
|
||||
.await
|
||||
.ok_or_else(|| Error::RouteError(Some("route peer info not found".to_string())))?;
|
||||
if info.noise_static_pubkey.is_empty() {
|
||||
return Err(Error::RouteError(Some(
|
||||
"remote static pubkey not found".to_string(),
|
||||
)));
|
||||
}
|
||||
Ok(info.noise_static_pubkey)
|
||||
}
|
||||
|
||||
fn get_handshake_lock(&self, peer_id: PeerId) -> Arc<Mutex<()>> {
|
||||
self.handshake_locks
|
||||
.entry(peer_id)
|
||||
.or_insert_with(|| Arc::new(Mutex::new(())))
|
||||
.clone()
|
||||
}
|
||||
|
||||
async fn send_handshake_packet(
|
||||
&self,
|
||||
payload: Vec<u8>,
|
||||
packet_type: PacketType,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<(), Error> {
|
||||
let mut pkt = ZCPacket::new_with_payload(&payload);
|
||||
pkt.fill_peer_manager_hdr(self.my_peer_id, dst_peer_id, packet_type as u8);
|
||||
let pkt_len = pkt.buf_len() as u64;
|
||||
self.send_via_next_hop(pkt, dst_peer_id, policy).await?;
|
||||
self.context
|
||||
.record_control_tx(&self.metric_network_name, pkt_len);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn send_via_next_hop(
|
||||
&self,
|
||||
msg: ZCPacket,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<(), Error> {
|
||||
self.route_transport
|
||||
.send_msg_to_next_hop(msg, dst_peer_id, policy)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn send_msg(
|
||||
self: &Arc<Self>,
|
||||
mut msg: ZCPacket,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<(), Error> {
|
||||
let now = Instant::now();
|
||||
|
||||
self.states.entry(dst_peer_id).or_default().last_active_at = now;
|
||||
|
||||
if self.is_secure_mode_enabled() {
|
||||
match self.ensure_session(dst_peer_id, policy.clone()).await {
|
||||
Ok(session) => {
|
||||
let my_peer_id = self.my_peer_id;
|
||||
session
|
||||
.encrypt_payload(my_peer_id, dst_peer_id, &mut msg)
|
||||
.map_err(|e| Error::RouteError(Some(format!("{e:?}"))))?;
|
||||
}
|
||||
Err(_) => {
|
||||
// Handshake in progress, buffer the packet instead of dropping it
|
||||
self.buffer_pending_packet(dst_peer_id, msg, policy);
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.send_via_next_hop(msg, dst_peer_id, policy).await
|
||||
}
|
||||
|
||||
fn buffer_pending_packet(&self, dst_peer_id: PeerId, pkt: ZCPacket, policy: NextHopPolicy) {
|
||||
let mut entry = self.pending_packets.entry(dst_peer_id).or_default();
|
||||
if entry.len() < MAX_PENDING_PACKETS_PER_PEER {
|
||||
entry.push((pkt, policy));
|
||||
}
|
||||
// silently drop when buffer is full
|
||||
}
|
||||
|
||||
async fn flush_pending_packets(&self, dst_peer_id: PeerId, session: Arc<PeerSession>) {
|
||||
let packets = self.pending_packets.remove(&dst_peer_id).map(|(_, v)| v);
|
||||
let Some(packets) = packets else { return };
|
||||
if packets.is_empty() {
|
||||
return;
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
?dst_peer_id,
|
||||
count = packets.len(),
|
||||
"flushing pending packets after relay handshake"
|
||||
);
|
||||
|
||||
for (mut pkt, policy) in packets {
|
||||
if session
|
||||
.encrypt_payload(self.my_peer_id, dst_peer_id, &mut pkt)
|
||||
.is_err()
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let _ = self.send_via_next_hop(pkt, dst_peer_id, policy).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn has_session(&self, dst_peer_id: PeerId) -> bool {
|
||||
self.peer_session_store
|
||||
.get(&SessionKey::new(
|
||||
self.context.network_identity().network_name,
|
||||
dst_peer_id,
|
||||
))
|
||||
.is_some()
|
||||
}
|
||||
|
||||
pub async fn ensure_session(
|
||||
self: &Arc<Self>,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<Arc<PeerSession>, Error> {
|
||||
let network = self.context.network_identity();
|
||||
let key = SessionKey::new(network.network_name.clone(), dst_peer_id);
|
||||
if let Some(session) = self.peer_session_store.get(&key) {
|
||||
return Ok(session);
|
||||
}
|
||||
|
||||
let lock = self.get_handshake_lock(dst_peer_id);
|
||||
if let Ok(guard) = lock.try_lock_owned() {
|
||||
let self_clone = self.clone();
|
||||
tokio::spawn(async move {
|
||||
self_clone
|
||||
.handshake_session(dst_peer_id, policy, Some(guard))
|
||||
.await
|
||||
});
|
||||
};
|
||||
Err(Error::RouteError(Some(
|
||||
"relay handshake in progress".to_string(),
|
||||
)))
|
||||
}
|
||||
|
||||
#[tracing::instrument(skip(self, _lock_guard), level = "debug", ret)]
|
||||
pub async fn handshake_session(
|
||||
&self,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
_lock_guard: Option<OwnedMutexGuard<()>>,
|
||||
) -> Result<(), Error> {
|
||||
let network = self.context.network_identity();
|
||||
let key = SessionKey::new(network.network_name.clone(), dst_peer_id);
|
||||
if let Some(session) = self.peer_session_store.get(&key) {
|
||||
self.flush_pending_packets(dst_peer_id, session).await;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
if let Some(next_retry_at) = self.states.get(&dst_peer_id).and_then(|v| v.next_retry_at)
|
||||
&& Instant::now() < next_retry_at
|
||||
{
|
||||
self.pending_packets.remove(&dst_peer_id);
|
||||
return Err(Error::RouteError(Some(
|
||||
"relay handshake backoff".to_string(),
|
||||
)));
|
||||
}
|
||||
|
||||
let mut last_err = None;
|
||||
for attempt in 0..HANDSHAKE_MAX_ATTEMPTS {
|
||||
let ret = self
|
||||
.handshake_session_once(dst_peer_id, policy.clone())
|
||||
.await;
|
||||
match ret {
|
||||
Ok(session) => {
|
||||
self.register_handshake_success(dst_peer_id);
|
||||
self.flush_pending_packets(dst_peer_id, session).await;
|
||||
return Ok(());
|
||||
}
|
||||
Err(e) => {
|
||||
last_err = Some(e);
|
||||
self.register_handshake_failure(dst_peer_id, attempt);
|
||||
if attempt + 1 < HANDSHAKE_MAX_ATTEMPTS {
|
||||
let backoff = HANDSHAKE_RETRY_BASE_MS.saturating_mul(1 << attempt);
|
||||
crate::foundation::time::sleep(Duration::from_millis(backoff)).await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// All attempts failed, drop buffered packets
|
||||
self.pending_packets.remove(&dst_peer_id);
|
||||
|
||||
Err(last_err
|
||||
.unwrap_or_else(|| Error::RouteError(Some("relay handshake failed".to_string()))))
|
||||
}
|
||||
|
||||
#[tracing::instrument(skip(self), level = "debug", ret)]
|
||||
async fn handshake_session_once(
|
||||
&self,
|
||||
dst_peer_id: PeerId,
|
||||
policy: NextHopPolicy,
|
||||
) -> Result<Arc<PeerSession>, Error> {
|
||||
let network = self.context.network_identity();
|
||||
let session_key = SessionKey::new(network.network_name.clone(), dst_peer_id);
|
||||
let (local_private_key, _local_public_key) = self.get_local_keypair()?;
|
||||
let remote_static = self.get_remote_static_pubkey(dst_peer_id).await?;
|
||||
let params: NoiseParams = "Noise_IK_25519_ChaChaPoly_SHA256"
|
||||
.parse()
|
||||
.map_err(|e| Error::RouteError(Some(format!("parse noise params failed: {e:?}"))))?;
|
||||
|
||||
let builder = snow::Builder::new(params);
|
||||
let mut hs = builder
|
||||
.prologue(RELAY_NOISE_PROLOGUE)
|
||||
.map_err(|e| Error::RouteError(Some(format!("set prologue failed: {e:?}"))))?
|
||||
.local_private_key(&local_private_key)
|
||||
.map_err(|e| Error::RouteError(Some(format!("set local key failed: {e:?}"))))?
|
||||
.remote_public_key(&remote_static)
|
||||
.map_err(|e| Error::RouteError(Some(format!("set remote key failed: {e:?}"))))?
|
||||
.build_initiator()
|
||||
.map_err(|e| Error::RouteError(Some(format!("build initiator failed: {e:?}"))))?;
|
||||
|
||||
let a_session_generation = self
|
||||
.peer_session_store
|
||||
.get(&session_key)
|
||||
.map(|s| s.session_generation());
|
||||
let a_conn_id = uuid::Uuid::new_v4();
|
||||
let msg1_pb = RelayNoiseMsg1Pb {
|
||||
version: RELAY_NOISE_VERSION,
|
||||
a_session_generation,
|
||||
a_conn_id: Some(a_conn_id.into()),
|
||||
client_encryption_algorithm: self.context.flags().encryption_algorithm,
|
||||
};
|
||||
let payload = msg1_pb.encode_to_vec();
|
||||
let mut out = vec![0u8; 4096];
|
||||
let out_len = hs
|
||||
.write_message(&payload, &mut out)
|
||||
.map_err(|e| Error::RouteError(Some(format!("noise write msg1 failed: {e:?}"))))?;
|
||||
let (tx, rx) = oneshot::channel();
|
||||
self.pending_handshakes.insert(dst_peer_id, tx);
|
||||
|
||||
let send_res = self
|
||||
.send_handshake_packet(
|
||||
out[..out_len].to_vec(),
|
||||
PacketType::RelayHandshake,
|
||||
dst_peer_id,
|
||||
policy,
|
||||
)
|
||||
.await;
|
||||
|
||||
if send_res.is_err() {
|
||||
self.pending_handshakes.remove(&dst_peer_id);
|
||||
}
|
||||
send_res?;
|
||||
let msg2_pkt = match timeout(Duration::from_secs(HANDSHAKE_TIMEOUT_SECS), rx).await {
|
||||
Ok(Ok(pkt)) => pkt,
|
||||
Ok(Err(_)) => {
|
||||
self.pending_handshakes.remove(&dst_peer_id);
|
||||
return Err(Error::RouteError(Some(
|
||||
"relay handshake canceled".to_string(),
|
||||
)));
|
||||
}
|
||||
Err(_) => {
|
||||
self.pending_handshakes.remove(&dst_peer_id);
|
||||
return Err(Error::RouteError(Some(
|
||||
"relay handshake timeout".to_string(),
|
||||
)));
|
||||
}
|
||||
};
|
||||
|
||||
let msg2_pb = self.decode_handshake_message::<RelayNoiseMsg2Pb>(
|
||||
PacketType::RelayHandshakeAck,
|
||||
&mut hs,
|
||||
msg2_pkt,
|
||||
)?;
|
||||
if msg2_pb.a_conn_id_echo != Some(a_conn_id.into()) {
|
||||
return Err(Error::RouteError(Some(
|
||||
"relay msg2 conn_id_echo mismatch".to_string(),
|
||||
)));
|
||||
}
|
||||
|
||||
let action = PeerConnSessionActionPb::try_from(msg2_pb.action)
|
||||
.map_err(|_| Error::RouteError(Some("invalid session action".to_string())))?;
|
||||
let session_action = match action {
|
||||
PeerConnSessionActionPb::Join => PeerSessionAction::Join,
|
||||
PeerConnSessionActionPb::Sync => PeerSessionAction::Sync,
|
||||
PeerConnSessionActionPb::Create => PeerSessionAction::Create,
|
||||
};
|
||||
let remote_static_key = if remote_static.len() == 32 {
|
||||
let mut key = [0u8; 32];
|
||||
key.copy_from_slice(&remote_static);
|
||||
Some(key)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let root_key_bytes = msg2_pb
|
||||
.root_key_32
|
||||
.as_deref()
|
||||
.filter(|v| v.len() == 32)
|
||||
.map(|v| {
|
||||
let mut key_bytes = [0u8; 32];
|
||||
key_bytes.copy_from_slice(v);
|
||||
key_bytes
|
||||
});
|
||||
let algo = self.context.flags().encryption_algorithm;
|
||||
let session = self
|
||||
.peer_session_store
|
||||
.apply_initiator_action(
|
||||
&session_key,
|
||||
session_action,
|
||||
msg2_pb.b_session_generation,
|
||||
root_key_bytes,
|
||||
msg2_pb.initial_epoch,
|
||||
algo,
|
||||
msg2_pb.server_encryption_algorithm.clone(),
|
||||
remote_static_key,
|
||||
)
|
||||
.map_err(|e| Error::RouteError(Some(format!("{e:?}"))))?;
|
||||
|
||||
Ok(session)
|
||||
}
|
||||
|
||||
fn register_handshake_success(&self, dst_peer_id: PeerId) {
|
||||
let mut entry = self.states.entry(dst_peer_id).or_default();
|
||||
entry.failure_count = 0;
|
||||
entry.next_retry_at = None;
|
||||
}
|
||||
|
||||
fn register_handshake_failure(&self, dst_peer_id: PeerId, attempt: u32) {
|
||||
let mut entry = self.states.entry(dst_peer_id).or_default();
|
||||
entry.failure_count = entry.failure_count.saturating_add(1);
|
||||
let backoff = HANDSHAKE_RETRY_BASE_MS.saturating_mul(1 << attempt);
|
||||
entry.next_retry_at = Some(Instant::now() + Duration::from_millis(backoff));
|
||||
}
|
||||
|
||||
fn decode_handshake_message<MsgT: Message + Default>(
|
||||
&self,
|
||||
expected_type: PacketType,
|
||||
hs: &mut snow::HandshakeState,
|
||||
pkt: ZCPacket,
|
||||
) -> Result<MsgT, Error> {
|
||||
let hdr = pkt.peer_manager_header().ok_or_else(|| {
|
||||
Error::RouteError(Some("packet without peer manager header".to_string()))
|
||||
})?;
|
||||
if hdr.packet_type != expected_type as u8 {
|
||||
return Err(Error::RouteError(Some("packet type mismatch".to_string())));
|
||||
}
|
||||
let mut out = vec![0u8; 4096];
|
||||
let out_len = hs
|
||||
.read_message(pkt.payload(), &mut out)
|
||||
.map_err(|e| Error::RouteError(Some(format!("noise read msg failed: {e:?}"))))?;
|
||||
let msg = MsgT::decode(&out[..out_len])
|
||||
.map_err(|e| Error::RouteError(Some(format!("decode message failed: {e:?}"))))?;
|
||||
Ok(msg)
|
||||
}
|
||||
|
||||
pub async fn handle_handshake_packet(&self, packet: ZCPacket) -> Result<(), Error> {
|
||||
let hdr = packet
|
||||
.peer_manager_header()
|
||||
.ok_or_else(|| Error::RouteError(Some("packet without header".to_string())))?;
|
||||
let src_peer_id = hdr.from_peer_id.get();
|
||||
self.context
|
||||
.record_control_rx(&self.metric_network_name, packet.buf_len() as u64);
|
||||
match hdr.packet_type {
|
||||
x if x == PacketType::RelayHandshake as u8 => {
|
||||
tracing::debug!("handle_relay_msg1 from {:?}", src_peer_id);
|
||||
self.handle_relay_msg1(packet, src_peer_id).await
|
||||
}
|
||||
x if x == PacketType::RelayHandshakeAck as u8 => {
|
||||
if let Some((_, sender)) = self.pending_handshakes.remove(&src_peer_id) {
|
||||
let _ = sender.send(packet);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
_ => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_relay_msg1(&self, msg1: ZCPacket, remote_peer_id: PeerId) -> Result<(), Error> {
|
||||
// Check for bidirectional handshake race condition.
|
||||
// If we are also waiting for a RelayHandshakeAck from this peer,
|
||||
// use deterministic rule: the peer with smaller peer_id becomes initiator.
|
||||
if self.pending_handshakes.contains_key(&remote_peer_id) {
|
||||
// We have a pending handshake as initiator.
|
||||
// If remote_peer_id < my_peer_id, remote should be initiator, we should be responder.
|
||||
// Cancel our pending handshake and proceed as responder.
|
||||
if remote_peer_id < self.my_peer_id {
|
||||
tracing::debug!(
|
||||
?remote_peer_id,
|
||||
my_peer_id = ?self.my_peer_id,
|
||||
"bidirectional handshake race: yielding initiator role to smaller peer_id"
|
||||
);
|
||||
// Remove our pending handshake
|
||||
self.pending_handshakes.remove(&remote_peer_id);
|
||||
} else {
|
||||
// We have smaller peer_id, we should remain initiator.
|
||||
// Ignore this RelayHandshake and let our initiator flow complete.
|
||||
tracing::debug!(
|
||||
?remote_peer_id,
|
||||
my_peer_id = ?self.my_peer_id,
|
||||
"bidirectional handshake race: keeping initiator role due to smaller peer_id"
|
||||
);
|
||||
return Err(Error::RouteError(Some(
|
||||
"bidirectional handshake race: we are initiator".to_string(),
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
let (local_private_key, _local_public_key) = self.get_local_keypair()?;
|
||||
let params: NoiseParams = "Noise_IK_25519_ChaChaPoly_SHA256"
|
||||
.parse()
|
||||
.map_err(|e| Error::RouteError(Some(format!("parse noise params failed: {e:?}"))))?;
|
||||
let builder = snow::Builder::new(params);
|
||||
let mut hs = builder
|
||||
.prologue(RELAY_NOISE_PROLOGUE)
|
||||
.map_err(|e| Error::RouteError(Some(format!("set prologue failed: {e:?}"))))?
|
||||
.local_private_key(&local_private_key)
|
||||
.map_err(|e| Error::RouteError(Some(format!("set local key failed: {e:?}"))))?
|
||||
.build_responder()
|
||||
.map_err(|e| Error::RouteError(Some(format!("build responder failed: {e:?}"))))?;
|
||||
|
||||
let msg1_pb = self.decode_handshake_message::<RelayNoiseMsg1Pb>(
|
||||
PacketType::RelayHandshake,
|
||||
&mut hs,
|
||||
msg1,
|
||||
)?;
|
||||
let remote_static = hs
|
||||
.get_remote_static()
|
||||
.map(|x: &[u8]| x.to_vec())
|
||||
.unwrap_or_default();
|
||||
let remote_static_key = if remote_static.len() == 32 {
|
||||
let mut key = [0u8; 32];
|
||||
key.copy_from_slice(&remote_static);
|
||||
Some(key)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// Verify initiator's static public key matches the expected key from route info
|
||||
let expected_pubkey = self.get_remote_static_pubkey(remote_peer_id).await?;
|
||||
if remote_static != expected_pubkey {
|
||||
return Err(Error::RouteError(Some(format!(
|
||||
"responder: initiator static pubkey mismatch for peer {}, expected {} bytes, got {} bytes",
|
||||
remote_peer_id,
|
||||
expected_pubkey.len(),
|
||||
remote_static.len()
|
||||
))));
|
||||
}
|
||||
|
||||
let server_network_name = self.context.network_name();
|
||||
let algo = self.context.flags().encryption_algorithm;
|
||||
let key = SessionKey::new(server_network_name.clone(), remote_peer_id);
|
||||
let upsert = self
|
||||
.peer_session_store
|
||||
.upsert_responder_session(
|
||||
&key,
|
||||
msg1_pb.a_session_generation,
|
||||
algo.clone(),
|
||||
msg1_pb.client_encryption_algorithm.clone(),
|
||||
remote_static_key,
|
||||
)
|
||||
.map_err(|e| Error::RouteError(Some(format!("{e:?}"))))?;
|
||||
let msg2_pb = RelayNoiseMsg2Pb {
|
||||
action: match upsert.action {
|
||||
PeerSessionAction::Join => PeerConnSessionActionPb::Join as i32,
|
||||
PeerSessionAction::Sync => PeerConnSessionActionPb::Sync as i32,
|
||||
PeerSessionAction::Create => PeerConnSessionActionPb::Create as i32,
|
||||
},
|
||||
b_session_generation: upsert.session_generation,
|
||||
root_key_32: upsert.root_key.map(|k| k.to_vec()),
|
||||
initial_epoch: upsert.initial_epoch,
|
||||
b_conn_id: Some(uuid::Uuid::new_v4().into()),
|
||||
a_conn_id_echo: msg1_pb.a_conn_id,
|
||||
server_encryption_algorithm: algo,
|
||||
};
|
||||
let payload = msg2_pb.encode_to_vec();
|
||||
let mut out = vec![0u8; 4096];
|
||||
let out_len = hs
|
||||
.write_message(&payload, &mut out)
|
||||
.map_err(|e| Error::RouteError(Some(format!("noise write msg2 failed: {e:?}"))))?;
|
||||
|
||||
self.register_handshake_success(remote_peer_id);
|
||||
|
||||
self.send_handshake_packet(
|
||||
out[..out_len].to_vec(),
|
||||
PacketType::RelayHandshakeAck,
|
||||
remote_peer_id,
|
||||
NextHopPolicy::LeastHop,
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Flush any packets buffered while waiting for the handshake to complete
|
||||
self.flush_pending_packets(remote_peer_id, upsert.session)
|
||||
.await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn decrypt_if_needed(self: &Arc<Self>, packet: &mut ZCPacket) -> Result<bool, Error> {
|
||||
if !self.is_secure_mode_enabled() {
|
||||
return Ok(false);
|
||||
}
|
||||
let hdr = packet
|
||||
.peer_manager_header()
|
||||
.ok_or_else(|| Error::RouteError(Some("packet without header".to_string())))?;
|
||||
let from_peer_id = hdr.from_peer_id.get();
|
||||
let network = self.context.network_identity();
|
||||
let key = SessionKey::new(network.network_name.clone(), from_peer_id);
|
||||
let Some(session) = self.peer_session_store.get(&key) else {
|
||||
tracing::debug!(
|
||||
"relay session not found for peer {}, try handshake",
|
||||
from_peer_id
|
||||
);
|
||||
self.ensure_session(from_peer_id, NextHopPolicy::LeastHop)
|
||||
.await?;
|
||||
return Ok(false);
|
||||
};
|
||||
let now = Instant::now();
|
||||
let mut entry = self.states.entry(from_peer_id).or_default();
|
||||
entry.last_active_at = now;
|
||||
session.decrypt_payload(from_peer_id, self.my_peer_id, packet)?;
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
pub fn evict_idle_sessions(&self, idle: Duration) {
|
||||
let now = Instant::now();
|
||||
let mut to_remove = Vec::new();
|
||||
for entry in self.states.iter() {
|
||||
if now.duration_since(entry.last_active_at) > idle {
|
||||
to_remove.push(*entry.key());
|
||||
}
|
||||
}
|
||||
for peer_id in to_remove {
|
||||
self.states.remove(&peer_id);
|
||||
self.pending_handshakes.remove(&peer_id);
|
||||
self.handshake_locks.remove(&peer_id);
|
||||
self.pending_packets.remove(&peer_id);
|
||||
}
|
||||
shrink_dashmap(&self.states, None);
|
||||
shrink_dashmap(&self.pending_handshakes, None);
|
||||
shrink_dashmap(&self.handshake_locks, None);
|
||||
shrink_dashmap(&self.pending_packets, None);
|
||||
}
|
||||
|
||||
pub fn has_state(&self, peer_id: PeerId) -> bool {
|
||||
self.states.contains_key(&peer_id)
|
||||
}
|
||||
|
||||
/// Remove relay-specific state for a specific peer.
|
||||
/// This does NOT remove the session from PeerSessionStore, because the
|
||||
/// session lifecycle is independent of any particular connection type
|
||||
/// (relay or direct). The session may still be used by direct connections
|
||||
/// or for fast reconnection (Join instead of Create).
|
||||
pub fn remove_peer(&self, peer_id: PeerId) {
|
||||
self.states.remove(&peer_id);
|
||||
self.pending_handshakes.remove(&peer_id);
|
||||
self.handshake_locks.remove(&peer_id);
|
||||
self.pending_packets.remove(&peer_id);
|
||||
shrink_dashmap(&self.states, None);
|
||||
shrink_dashmap(&self.pending_handshakes, None);
|
||||
shrink_dashmap(&self.handshake_locks, None);
|
||||
shrink_dashmap(&self.pending_packets, None);
|
||||
|
||||
tracing::debug!(?peer_id, "RelayPeerMap removed peer relay state");
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(any(test, feature = "test-utils"))]
|
||||
mod test_utils {
|
||||
use super::*;
|
||||
|
||||
impl RelayPeerMap {
|
||||
#[doc(hidden)]
|
||||
pub(crate) fn has_session_without_touch(&self, dst_peer_id: PeerId) -> bool {
|
||||
self.peer_session_store.contains_valid(&SessionKey::new(
|
||||
self.context.network_identity().network_name,
|
||||
dst_peer_id,
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
use core::cmp::Ordering;
|
||||
use petgraph::{
|
||||
algo::Measure,
|
||||
visit::{EdgeRef, IntoEdges, VisitMap, Visitable},
|
||||
};
|
||||
use std::collections::HashMap;
|
||||
use std::collections::hash_map::Entry::{Occupied, Vacant};
|
||||
use std::{collections::BinaryHeap, hash::Hash};
|
||||
|
||||
/// `MinScored<K, T>` holds a score `K` and a scored object `T` in
|
||||
/// a pair for use with a `BinaryHeap`.
|
||||
///
|
||||
/// `MinScored` compares in reverse order by the score, so that we can
|
||||
/// use `BinaryHeap` as a min-heap to extract the score-value pair with the
|
||||
/// least score.
|
||||
///
|
||||
/// **Note:** `MinScored` implements a total order (`Ord`), so that it is
|
||||
/// possible to use float types as scores.
|
||||
#[derive(Copy, Clone, Debug)]
|
||||
pub struct MinScored<K, T>(pub K, pub T);
|
||||
|
||||
impl<K: PartialOrd, T> PartialEq for MinScored<K, T> {
|
||||
#[inline]
|
||||
fn eq(&self, other: &MinScored<K, T>) -> bool {
|
||||
self.cmp(other) == Ordering::Equal
|
||||
}
|
||||
}
|
||||
|
||||
impl<K: PartialOrd, T> Eq for MinScored<K, T> {}
|
||||
|
||||
impl<K: PartialOrd, T> PartialOrd for MinScored<K, T> {
|
||||
#[inline]
|
||||
fn partial_cmp(&self, other: &MinScored<K, T>) -> Option<Ordering> {
|
||||
Some(self.cmp(other))
|
||||
}
|
||||
}
|
||||
|
||||
impl<K: PartialOrd, T> Ord for MinScored<K, T> {
|
||||
#[inline]
|
||||
fn cmp(&self, other: &MinScored<K, T>) -> Ordering {
|
||||
let a = &self.0;
|
||||
let b = &other.0;
|
||||
if a == b {
|
||||
Ordering::Equal
|
||||
} else if a < b {
|
||||
Ordering::Greater
|
||||
} else if a > b {
|
||||
Ordering::Less
|
||||
} else if a.ne(a) && b.ne(b) {
|
||||
// these are the NaN cases
|
||||
Ordering::Equal
|
||||
} else if a.ne(a) {
|
||||
// Order NaN less, so that it is last in the MinScore order
|
||||
Ordering::Less
|
||||
} else {
|
||||
Ordering::Greater
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub type DijkstraResult<K, NodeId> = (HashMap<NodeId, K>, HashMap<NodeId, (NodeId, usize)>);
|
||||
|
||||
pub fn dijkstra_with_first_hop<G, F, K>(
|
||||
graph: G,
|
||||
start: G::NodeId,
|
||||
mut edge_cost: F,
|
||||
) -> DijkstraResult<K, G::NodeId>
|
||||
where
|
||||
G: IntoEdges + Visitable,
|
||||
G::NodeId: Eq + Hash + Clone,
|
||||
F: FnMut(G::EdgeRef) -> K,
|
||||
K: Measure + Copy,
|
||||
{
|
||||
let mut visited = graph.visit_map();
|
||||
let mut scores = HashMap::new();
|
||||
let mut first_hop = HashMap::new();
|
||||
let mut visit_next = BinaryHeap::new();
|
||||
let zero_score = K::default();
|
||||
scores.insert(start, zero_score);
|
||||
visit_next.push(MinScored(zero_score, start));
|
||||
first_hop.insert(start, (start, 0));
|
||||
|
||||
while let Some(MinScored(node_score, node)) = visit_next.pop() {
|
||||
if visited.is_visited(&node) {
|
||||
continue;
|
||||
}
|
||||
for edge in graph.edges(node) {
|
||||
let next = edge.target();
|
||||
if visited.is_visited(&next) {
|
||||
continue;
|
||||
}
|
||||
let next_score = node_score + edge_cost(edge);
|
||||
match scores.entry(next) {
|
||||
Occupied(mut ent) => {
|
||||
if next_score < *ent.get() {
|
||||
*ent.get_mut() = next_score;
|
||||
visit_next.push(MinScored(next_score, next));
|
||||
// 继承前驱的 first_hop,或自己就是第一跳
|
||||
let hop = if node == start {
|
||||
(next, 0)
|
||||
} else {
|
||||
first_hop[&node]
|
||||
};
|
||||
first_hop.insert(next, (hop.0, hop.1 + 1));
|
||||
}
|
||||
}
|
||||
Vacant(ent) => {
|
||||
ent.insert(next_score);
|
||||
visit_next.push(MinScored(next_score, next));
|
||||
let hop = if node == start {
|
||||
(next, 0)
|
||||
} else {
|
||||
first_hop[&node]
|
||||
};
|
||||
first_hop.insert(next, (hop.0, hop.1 + 1));
|
||||
}
|
||||
}
|
||||
}
|
||||
visited.visit(node);
|
||||
}
|
||||
|
||||
(scores, first_hop)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use petgraph::graph::DiGraph;
|
||||
|
||||
#[test]
|
||||
fn test_dijkstra_with_first_hop_4node() {
|
||||
let mut graph = DiGraph::<&str, u32>::new();
|
||||
let a = graph.add_node("a");
|
||||
let b = graph.add_node("b");
|
||||
let c = graph.add_node("c");
|
||||
let d = graph.add_node("d");
|
||||
|
||||
graph.extend_with_edges([(a, b, 1)]);
|
||||
graph.extend_with_edges([(b, c, 1)]);
|
||||
graph.extend_with_edges([(c, d, 2)]);
|
||||
|
||||
let (scores, first_hop) = dijkstra_with_first_hop(&graph, a, |edge| *edge.weight());
|
||||
|
||||
assert_eq!(scores[&b], 1);
|
||||
assert_eq!(scores[&c], 2);
|
||||
assert_eq!(scores[&d], 4);
|
||||
|
||||
assert_eq!(first_hop[&b], (b, 1));
|
||||
assert_eq!(first_hop[&c], (b, 2));
|
||||
assert_eq!(first_hop[&d], (b, 3));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_dijkstra_with_first_hop() {
|
||||
let mut graph = DiGraph::<&str, u32>::new();
|
||||
let a = graph.add_node("a");
|
||||
let b = graph.add_node("b");
|
||||
let c = graph.add_node("c");
|
||||
let d = graph.add_node("d");
|
||||
let e = graph.add_node("e");
|
||||
|
||||
graph.extend_with_edges([(a, b, 1), (a, c, 2), (b, d, 1), (c, d, 3), (d, e, 1)]);
|
||||
|
||||
let (scores, first_hop) = dijkstra_with_first_hop(&graph, a, |edge| *edge.weight());
|
||||
|
||||
assert_eq!(scores[&b], 1);
|
||||
assert_eq!(scores[&c], 2);
|
||||
assert_eq!(scores[&d], 2);
|
||||
assert_eq!(scores[&e], 3);
|
||||
|
||||
assert_eq!(first_hop[&b], (b, 1));
|
||||
assert_eq!(first_hop[&c], (c, 1));
|
||||
assert_eq!(first_hop[&d], (b, 2)); // d is reached via b
|
||||
assert_eq!(first_hop[&e], (b, 3)); // e is reached via d
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
//! Route trait surface shared by the peer domain, plus the OSPF route
|
||||
//! implementation and the graph algorithms backing it.
|
||||
|
||||
pub(crate) mod graph_algo;
|
||||
pub(crate) mod peer_ospf_route;
|
||||
mod route_peer_wire;
|
||||
|
||||
use cidr::Ipv6Inet;
|
||||
use cidr::{Ipv4Cidr, Ipv6Cidr};
|
||||
use dashmap::DashMap;
|
||||
use quanta::Instant;
|
||||
use std::{
|
||||
collections::BTreeSet,
|
||||
net::{Ipv4Addr, Ipv6Addr},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use crate::{
|
||||
config::PeerId,
|
||||
peers::context::NetworkIdentity,
|
||||
proto::{
|
||||
core_peer::peer::{ListPublicIpv6InfoResponse, Route as CoreRoute},
|
||||
peer_rpc::{
|
||||
ForeignNetworkRouteInfoEntry, ForeignNetworkRouteInfoKey, PeerIdentityType,
|
||||
RouteForeignNetworkInfos, RouteForeignNetworkSummary, RoutePeerInfo,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub enum NextHopPolicy {
|
||||
#[default]
|
||||
LeastHop,
|
||||
LeastCost,
|
||||
}
|
||||
|
||||
pub type ForeignNetworkRouteInfoMap =
|
||||
DashMap<ForeignNetworkRouteInfoKey, ForeignNetworkRouteInfoEntry>;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
pub trait RouteInterface {
|
||||
async fn list_peers(&self) -> Vec<PeerId>;
|
||||
fn my_peer_id(&self) -> PeerId;
|
||||
fn need_periodic_requery_peers(&self) -> bool {
|
||||
false
|
||||
}
|
||||
async fn close_peer(&self, _peer_id: PeerId) {}
|
||||
async fn get_peer_public_key(&self, _peer_id: PeerId) -> Option<Vec<u8>> {
|
||||
None
|
||||
}
|
||||
async fn get_peer_identity_type(&self, _peer_id: PeerId) -> Option<PeerIdentityType> {
|
||||
None
|
||||
}
|
||||
async fn list_foreign_networks(&self) -> ForeignNetworkRouteInfoMap {
|
||||
DashMap::new()
|
||||
}
|
||||
}
|
||||
|
||||
pub type RouteInterfaceBox = Box<dyn RouteInterface + Send + Sync>;
|
||||
|
||||
#[auto_impl::auto_impl(Box , &mut)]
|
||||
pub trait RouteCostCalculatorInterface: Send + Sync {
|
||||
fn begin_update(&mut self) {}
|
||||
fn end_update(&mut self) {}
|
||||
|
||||
fn calculate_cost(&self, _src: PeerId, _dst: PeerId) -> i32 {
|
||||
1
|
||||
}
|
||||
|
||||
fn need_update(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn dump(&self) -> String {
|
||||
"All routes have cost 1".to_string()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub struct DefaultRouteCostCalculator;
|
||||
|
||||
impl RouteCostCalculatorInterface for DefaultRouteCostCalculator {}
|
||||
|
||||
pub type RouteCostCalculator = Box<dyn RouteCostCalculatorInterface>;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
#[auto_impl::auto_impl(Box, Arc)]
|
||||
pub trait Route {
|
||||
async fn open(&self, interface: RouteInterfaceBox) -> Result<u8, ()>;
|
||||
async fn close(&self);
|
||||
|
||||
async fn get_next_hop(&self, peer_id: PeerId) -> Option<PeerId>;
|
||||
async fn get_next_hop_with_policy(
|
||||
&self,
|
||||
peer_id: PeerId,
|
||||
_policy: NextHopPolicy,
|
||||
) -> Option<PeerId> {
|
||||
self.get_next_hop(peer_id).await
|
||||
}
|
||||
|
||||
async fn list_routes(&self) -> Vec<CoreRoute>;
|
||||
|
||||
// TODO: rewrite route management, remove this
|
||||
async fn list_proxy_cidrs(&self) -> BTreeSet<Ipv4Cidr>;
|
||||
|
||||
// TODO: rewrite route management, remove this
|
||||
async fn list_proxy_cidrs_v6(&self) -> BTreeSet<Ipv6Cidr>;
|
||||
|
||||
async fn list_public_ipv6_routes(&self) -> BTreeSet<Ipv6Inet> {
|
||||
BTreeSet::new()
|
||||
}
|
||||
|
||||
async fn get_my_public_ipv6_addr(&self) -> Option<Ipv6Inet> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn get_public_ipv6_gateway_peer_id(&self) -> Option<PeerId> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn get_local_public_ipv6_info(&self) -> ListPublicIpv6InfoResponse {
|
||||
ListPublicIpv6InfoResponse::default()
|
||||
}
|
||||
|
||||
async fn get_peer_id_by_ipv4(&self, _ipv4: &Ipv4Addr) -> Option<PeerId> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn get_peer_id_by_ipv6(&self, _ipv6: &Ipv6Addr) -> Option<PeerId> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn get_peer_id_by_ip(&self, ip: &std::net::IpAddr) -> Option<PeerId> {
|
||||
match ip {
|
||||
std::net::IpAddr::V4(v4) => self.get_peer_id_by_ipv4(v4).await,
|
||||
std::net::IpAddr::V6(v6) => self.get_peer_id_by_ipv6(v6).await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_peers_own_foreign_network(
|
||||
&self,
|
||||
_network_identity: &NetworkIdentity,
|
||||
) -> Vec<PeerId> {
|
||||
vec![]
|
||||
}
|
||||
|
||||
async fn list_foreign_network_info(&self) -> RouteForeignNetworkInfos {
|
||||
Default::default()
|
||||
}
|
||||
|
||||
async fn get_foreign_network_summary(&self) -> RouteForeignNetworkSummary {
|
||||
Default::default()
|
||||
}
|
||||
|
||||
// my peer id in foreign network is different from the one in local network
|
||||
// this function is used to get the peer id in local network
|
||||
async fn get_origin_my_peer_id(
|
||||
&self,
|
||||
_network_name: &str,
|
||||
_foreign_my_peer_id: PeerId,
|
||||
) -> Option<PeerId> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn set_route_cost_fn(&self, _cost_fn: RouteCostCalculator) {}
|
||||
|
||||
async fn get_peer_info(&self, peer_id: PeerId) -> Option<RoutePeerInfo>;
|
||||
|
||||
async fn get_peer_info_last_update_time(&self) -> Instant;
|
||||
|
||||
fn get_peer_groups(&self, peer_id: PeerId) -> Arc<Vec<String>>;
|
||||
|
||||
async fn refresh_acl_groups(&self) {}
|
||||
|
||||
async fn get_peer_groups_by_ip(&self, ip: &std::net::IpAddr) -> Arc<Vec<String>> {
|
||||
match self.get_peer_id_by_ip(ip).await {
|
||||
Some(peer_id) => self.get_peer_groups(peer_id),
|
||||
None => Arc::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
async fn dump(&self) -> String {
|
||||
"this route implementation does not support dump".to_string()
|
||||
}
|
||||
}
|
||||
|
||||
pub type ArcRoute = Arc<dyn Route + Send + Sync>;
|
||||
|
||||
pub(crate) struct DisabledRoute;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Route for DisabledRoute {
|
||||
async fn open(&self, _interface: RouteInterfaceBox) -> Result<u8, ()> {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
async fn close(&self) {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
async fn get_next_hop(&self, _peer_id: PeerId) -> Option<PeerId> {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
async fn list_routes(&self) -> Vec<CoreRoute> {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
// TODO: rewrite route management, remove this
|
||||
async fn list_proxy_cidrs(&self) -> BTreeSet<Ipv4Cidr> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
// TODO: rewrite route management, remove this
|
||||
async fn list_proxy_cidrs_v6(&self) -> BTreeSet<Ipv6Cidr> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn list_public_ipv6_routes(&self) -> BTreeSet<Ipv6Inet> {
|
||||
unimplemented!()
|
||||
}
|
||||
|
||||
async fn get_my_public_ipv6_addr(&self) -> Option<Ipv6Inet> {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
async fn get_peer_info(&self, _peer_id: PeerId) -> Option<RoutePeerInfo> {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
async fn get_peer_info_last_update_time(&self) -> Instant {
|
||||
panic!("mock route")
|
||||
}
|
||||
|
||||
fn get_peer_groups(&self, _peer_id: PeerId) -> Arc<Vec<String>> {
|
||||
panic!("mock route")
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,399 @@
|
||||
//! Minimal protobuf wire editing used by OSPF route reflection.
|
||||
//!
|
||||
//! Route calculation uses generated prost types. This module only keeps the
|
||||
//! original `RoutePeerInfo` bytes and replaces the two fields that credential
|
||||
//! filtering is allowed to change, leaving every other field byte-for-byte
|
||||
//! intact.
|
||||
|
||||
use bytes::Bytes;
|
||||
use prost::{
|
||||
Message,
|
||||
encoding::{
|
||||
DecodeContext, WireType, decode_key, decode_varint, encode_key, encode_varint, skip_field,
|
||||
},
|
||||
};
|
||||
use thiserror::Error;
|
||||
|
||||
use crate::proto::peer_rpc::{RoutePeerInfo, SyncRouteInfoRequest};
|
||||
|
||||
const SYNC_ROUTE_PEER_INFOS_TAG: u32 = 4;
|
||||
const ROUTE_PEER_INFOS_ITEM_TAG: u32 = 1;
|
||||
const ROUTE_PEER_INFO_PROXY_CIDRS_TAG: u32 = 5;
|
||||
const ROUTE_PEER_INFO_FEATURE_FLAG_TAG: u32 = 11;
|
||||
const ROUTE_PEER_INFO_CREDENTIAL_PROOF_TAG: u32 = 19;
|
||||
const FEATURE_FLAG_IS_CREDENTIAL_PEER_TAG: u32 = 8;
|
||||
const CREDENTIAL_PROOF_CREDENTIAL_TAG: u32 = 1;
|
||||
|
||||
pub(crate) type RawRoutePeerInfo = Bytes;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub(crate) enum WireError {
|
||||
#[error(transparent)]
|
||||
Decode(#[from] prost::DecodeError),
|
||||
#[error("protobuf field {tag} has wire type {actual:?}, expected {expected:?}")]
|
||||
WrongWireType {
|
||||
tag: u32,
|
||||
actual: WireType,
|
||||
expected: WireType,
|
||||
},
|
||||
#[error("protobuf length-delimited field is larger than the remaining input")]
|
||||
TruncatedLengthDelimited,
|
||||
#[error("raw RoutePeerInfo count does not match the decoded request")]
|
||||
PeerInfoCountMismatch,
|
||||
}
|
||||
|
||||
type Result<T> = std::result::Result<T, WireError>;
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct WireField<'a> {
|
||||
tag: u32,
|
||||
wire_type: WireType,
|
||||
encoded: &'a [u8],
|
||||
length_delimited: Option<&'a [u8]>,
|
||||
}
|
||||
|
||||
fn parse_fields(message: &[u8]) -> Result<Vec<WireField<'_>>> {
|
||||
let mut input = message;
|
||||
let mut fields = Vec::new();
|
||||
|
||||
while !input.is_empty() {
|
||||
let start = message.len() - input.len();
|
||||
let (tag, wire_type) = decode_key(&mut input)?;
|
||||
let length_delimited = if wire_type == WireType::LengthDelimited {
|
||||
let len = usize::try_from(decode_varint(&mut input)?)
|
||||
.map_err(|_| WireError::TruncatedLengthDelimited)?;
|
||||
if len > input.len() {
|
||||
return Err(WireError::TruncatedLengthDelimited);
|
||||
}
|
||||
let (payload, remaining) = input.split_at(len);
|
||||
input = remaining;
|
||||
Some(payload)
|
||||
} else {
|
||||
skip_field(wire_type, tag, &mut input, DecodeContext::default())?;
|
||||
None
|
||||
};
|
||||
let end = message.len() - input.len();
|
||||
fields.push(WireField {
|
||||
tag,
|
||||
wire_type,
|
||||
encoded: &message[start..end],
|
||||
length_delimited,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(fields)
|
||||
}
|
||||
|
||||
fn length_delimited_fields(message: &[u8], tag: u32) -> Result<Vec<&[u8]>> {
|
||||
parse_fields(message)?
|
||||
.into_iter()
|
||||
.filter(|field| field.tag == tag)
|
||||
.map(|field| {
|
||||
field.length_delimited.ok_or(WireError::WrongWireType {
|
||||
tag,
|
||||
actual: field.wire_type,
|
||||
expected: WireType::LengthDelimited,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn replace_fields(
|
||||
message: &[u8],
|
||||
tag: u32,
|
||||
expected_wire_type: WireType,
|
||||
replacement: &[u8],
|
||||
) -> Result<Vec<u8>> {
|
||||
let mut output = Vec::with_capacity(message.len() + replacement.len());
|
||||
for field in parse_fields(message)? {
|
||||
if field.tag != tag {
|
||||
output.extend_from_slice(field.encoded);
|
||||
continue;
|
||||
}
|
||||
if field.wire_type != expected_wire_type {
|
||||
return Err(WireError::WrongWireType {
|
||||
tag,
|
||||
actual: field.wire_type,
|
||||
expected: expected_wire_type,
|
||||
});
|
||||
}
|
||||
}
|
||||
output.extend_from_slice(replacement);
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn append_length_delimited_field(output: &mut Vec<u8>, tag: u32, payload: &[u8]) {
|
||||
encode_key(tag, WireType::LengthDelimited, output);
|
||||
encode_varint(payload.len() as u64, output);
|
||||
output.extend_from_slice(payload);
|
||||
}
|
||||
|
||||
fn append_varint_field(output: &mut Vec<u8>, tag: u32, value: u64) {
|
||||
encode_key(tag, WireType::Varint, output);
|
||||
encode_varint(value, output);
|
||||
}
|
||||
|
||||
fn merged_length_delimited_field(message: &[u8], tag: u32) -> Result<Vec<u8>> {
|
||||
let values = length_delimited_fields(message, tag)?;
|
||||
let total_len = values.iter().map(|value| value.len()).sum();
|
||||
let mut merged = Vec::with_capacity(total_len);
|
||||
for value in values {
|
||||
merged.extend_from_slice(value);
|
||||
}
|
||||
Ok(merged)
|
||||
}
|
||||
|
||||
pub(crate) fn raw_route_peer_info(info: &RoutePeerInfo) -> RawRoutePeerInfo {
|
||||
Bytes::from(info.encode_to_vec())
|
||||
}
|
||||
|
||||
pub(crate) fn extract_route_peer_infos(request: &[u8]) -> Result<Vec<RawRoutePeerInfo>> {
|
||||
let mut result = Vec::new();
|
||||
for peer_infos in length_delimited_fields(request, SYNC_ROUTE_PEER_INFOS_TAG)? {
|
||||
result.extend(
|
||||
length_delimited_fields(peer_infos, ROUTE_PEER_INFOS_ITEM_TAG)?
|
||||
.into_iter()
|
||||
.map(Bytes::copy_from_slice),
|
||||
);
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
pub(crate) fn encode_sync_route_request(
|
||||
request: &SyncRouteInfoRequest,
|
||||
raw_peer_infos: &[RawRoutePeerInfo],
|
||||
) -> Result<Vec<u8>> {
|
||||
let decoded_count = request
|
||||
.peer_infos
|
||||
.as_ref()
|
||||
.map(|peer_infos| peer_infos.items.len())
|
||||
.unwrap_or_default();
|
||||
if decoded_count != raw_peer_infos.len() {
|
||||
return Err(WireError::PeerInfoCountMismatch);
|
||||
}
|
||||
|
||||
let mut request_without_peer_infos = request.clone();
|
||||
request_without_peer_infos.peer_infos = None;
|
||||
let mut output = request_without_peer_infos.encode_to_vec();
|
||||
if request.peer_infos.is_some() {
|
||||
let mut peer_infos = Vec::new();
|
||||
for info in raw_peer_infos {
|
||||
append_length_delimited_field(&mut peer_infos, ROUTE_PEER_INFOS_ITEM_TAG, info);
|
||||
}
|
||||
append_length_delimited_field(&mut output, SYNC_ROUTE_PEER_INFOS_TAG, &peer_infos);
|
||||
}
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
pub(crate) fn patch_credential_route_peer_info(
|
||||
raw: &RawRoutePeerInfo,
|
||||
proxy_cidrs: &[String],
|
||||
) -> Result<RawRoutePeerInfo> {
|
||||
let mut proxy_cidr_fields = Vec::new();
|
||||
for cidr in proxy_cidrs {
|
||||
append_length_delimited_field(
|
||||
&mut proxy_cidr_fields,
|
||||
ROUTE_PEER_INFO_PROXY_CIDRS_TAG,
|
||||
cidr.as_bytes(),
|
||||
);
|
||||
}
|
||||
let route_info = replace_fields(
|
||||
raw,
|
||||
ROUTE_PEER_INFO_PROXY_CIDRS_TAG,
|
||||
WireType::LengthDelimited,
|
||||
&proxy_cidr_fields,
|
||||
)?;
|
||||
|
||||
// Singular message fields merge when they occur more than once. Concatenating
|
||||
// their payloads preserves that protobuf behavior before changing tag 8.
|
||||
let feature_flag =
|
||||
merged_length_delimited_field(&route_info, ROUTE_PEER_INFO_FEATURE_FLAG_TAG)?;
|
||||
let mut credential_flag = Vec::new();
|
||||
append_varint_field(&mut credential_flag, FEATURE_FLAG_IS_CREDENTIAL_PEER_TAG, 1);
|
||||
let feature_flag = replace_fields(
|
||||
&feature_flag,
|
||||
FEATURE_FLAG_IS_CREDENTIAL_PEER_TAG,
|
||||
WireType::Varint,
|
||||
&credential_flag,
|
||||
)?;
|
||||
let mut feature_flag_field = Vec::new();
|
||||
append_length_delimited_field(
|
||||
&mut feature_flag_field,
|
||||
ROUTE_PEER_INFO_FEATURE_FLAG_TAG,
|
||||
&feature_flag,
|
||||
);
|
||||
|
||||
Ok(Bytes::from(replace_fields(
|
||||
&route_info,
|
||||
ROUTE_PEER_INFO_FEATURE_FLAG_TAG,
|
||||
WireType::LengthDelimited,
|
||||
&feature_flag_field,
|
||||
)?))
|
||||
}
|
||||
|
||||
pub(crate) fn raw_credential_bytes(
|
||||
raw_route_info: &RawRoutePeerInfo,
|
||||
proof_idx: usize,
|
||||
) -> Result<Option<Bytes>> {
|
||||
let proofs = length_delimited_fields(raw_route_info, ROUTE_PEER_INFO_CREDENTIAL_PROOF_TAG)?;
|
||||
let Some(proof) = proofs.get(proof_idx) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let credentials = length_delimited_fields(proof, CREDENTIAL_PROOF_CREDENTIAL_TAG)?;
|
||||
if credentials.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
// Multiple occurrences of a singular message field merge. Concatenation is
|
||||
// the wire-equivalent merged message and keeps nested unknown fields intact.
|
||||
let total_len = credentials.iter().map(|value| value.len()).sum();
|
||||
let mut merged = Vec::with_capacity(total_len);
|
||||
for credential in credentials {
|
||||
merged.extend_from_slice(credential);
|
||||
}
|
||||
Ok(Some(Bytes::from(merged)))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::proto::{
|
||||
common::PeerFeatureFlag,
|
||||
peer_rpc::{RoutePeerInfos, TrustedCredentialPubkey, TrustedCredentialPubkeyProof},
|
||||
};
|
||||
|
||||
fn encoded_varint_field(tag: u32, value: u64) -> Vec<u8> {
|
||||
let mut output = Vec::new();
|
||||
append_varint_field(&mut output, tag, value);
|
||||
output
|
||||
}
|
||||
|
||||
fn encoded_length_delimited_field(tag: u32, payload: &[u8]) -> Vec<u8> {
|
||||
let mut output = Vec::new();
|
||||
append_length_delimited_field(&mut output, tag, payload);
|
||||
output
|
||||
}
|
||||
|
||||
fn encoded_fields(message: &[u8], tag: u32) -> Vec<Vec<u8>> {
|
||||
parse_fields(message)
|
||||
.unwrap()
|
||||
.into_iter()
|
||||
.filter(|field| field.tag == tag)
|
||||
.map(|field| field.encoded.to_vec())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn patch_preserves_top_level_and_nested_unknown_fields() {
|
||||
const TOP_LEVEL_UNKNOWN_TAG: u32 = 100;
|
||||
const FEATURE_UNKNOWN_TAG: u32 = 101;
|
||||
|
||||
let mut feature_flag = PeerFeatureFlag {
|
||||
avoid_relay_data: true,
|
||||
..Default::default()
|
||||
}
|
||||
.encode_to_vec();
|
||||
feature_flag.extend(encoded_varint_field(FEATURE_UNKNOWN_TAG, 42));
|
||||
|
||||
let info = RoutePeerInfo {
|
||||
peer_id: 7,
|
||||
proxy_cidrs: vec!["10.0.0.0/8".to_owned()],
|
||||
..Default::default()
|
||||
};
|
||||
let mut raw = info.encode_to_vec();
|
||||
raw.extend(encoded_length_delimited_field(
|
||||
ROUTE_PEER_INFO_FEATURE_FLAG_TAG,
|
||||
&feature_flag,
|
||||
));
|
||||
raw.extend(encoded_varint_field(TOP_LEVEL_UNKNOWN_TAG, 99));
|
||||
|
||||
let top_level_unknown = encoded_fields(&raw, TOP_LEVEL_UNKNOWN_TAG);
|
||||
let nested_unknown = encoded_fields(&feature_flag, FEATURE_UNKNOWN_TAG);
|
||||
let patched =
|
||||
patch_credential_route_peer_info(&Bytes::from(raw), &["10.1.0.0/16".to_owned()])
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
encoded_fields(&patched, TOP_LEVEL_UNKNOWN_TAG),
|
||||
top_level_unknown
|
||||
);
|
||||
let patched_feature =
|
||||
merged_length_delimited_field(&patched, ROUTE_PEER_INFO_FEATURE_FLAG_TAG).unwrap();
|
||||
assert_eq!(
|
||||
encoded_fields(&patched_feature, FEATURE_UNKNOWN_TAG),
|
||||
nested_unknown
|
||||
);
|
||||
|
||||
let decoded = RoutePeerInfo::decode(patched).unwrap();
|
||||
assert_eq!(decoded.proxy_cidrs, ["10.1.0.0/16"]);
|
||||
assert!(decoded.feature_flag.unwrap().is_credential_peer);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn multi_hop_sync_keeps_raw_peer_info_exactly() {
|
||||
let info = RoutePeerInfo {
|
||||
peer_id: 9,
|
||||
..Default::default()
|
||||
};
|
||||
let mut raw = info.encode_to_vec();
|
||||
raw.extend(encoded_length_delimited_field(120, b"future"));
|
||||
let raw = Bytes::from(raw);
|
||||
let request = SyncRouteInfoRequest {
|
||||
my_peer_id: 1,
|
||||
peer_infos: Some(RoutePeerInfos { items: vec![info] }),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let first_hop = encode_sync_route_request(&request, std::slice::from_ref(&raw)).unwrap();
|
||||
let first_hop_raw = extract_route_peer_infos(&first_hop).unwrap();
|
||||
assert_eq!(first_hop_raw.as_slice(), std::slice::from_ref(&raw));
|
||||
assert_eq!(
|
||||
SyncRouteInfoRequest::decode(first_hop.as_slice())
|
||||
.unwrap()
|
||||
.peer_infos,
|
||||
request.peer_infos
|
||||
);
|
||||
|
||||
let second_hop = encode_sync_route_request(&request, &first_hop_raw).unwrap();
|
||||
assert_eq!(extract_route_peer_infos(&second_hop).unwrap(), [raw]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn credential_hmac_uses_exact_nested_message_bytes() {
|
||||
let secret = "wire-test-secret";
|
||||
let credential = TrustedCredentialPubkey {
|
||||
pubkey: vec![3; 32],
|
||||
..Default::default()
|
||||
};
|
||||
let mut raw_credential = credential.encode_to_vec();
|
||||
raw_credential.extend(encoded_varint_field(100, 1234));
|
||||
let hmac = TrustedCredentialPubkeyProof::generate_credential_hmac_from_bytes(
|
||||
&raw_credential,
|
||||
secret,
|
||||
);
|
||||
|
||||
let mut raw_proof =
|
||||
encoded_length_delimited_field(CREDENTIAL_PROOF_CREDENTIAL_TAG, &raw_credential);
|
||||
raw_proof.extend(encoded_length_delimited_field(2, &hmac));
|
||||
let mut raw_route_info = RoutePeerInfo {
|
||||
peer_id: 11,
|
||||
..Default::default()
|
||||
}
|
||||
.encode_to_vec();
|
||||
raw_route_info.extend(encoded_length_delimited_field(
|
||||
ROUTE_PEER_INFO_CREDENTIAL_PROOF_TAG,
|
||||
&raw_proof,
|
||||
));
|
||||
|
||||
let extracted = raw_credential_bytes(&Bytes::from(raw_route_info), 0)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(extracted, raw_credential);
|
||||
|
||||
let proof = TrustedCredentialPubkeyProof {
|
||||
credential: Some(credential),
|
||||
credential_hmac: hmac,
|
||||
};
|
||||
assert!(proof.verify_credential_hmac_with_bytes(&extracted, secret));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,100 @@
|
||||
//! Test-only peer-context fakes shared by peer-domain unit tests
|
||||
//! (`peers::tests`, `peers::route::peer_ospf_route::tests`, and
|
||||
//! `context::tests`). Kept out of `context.rs` so the context unit tests and
|
||||
//! their consumers share one definition.
|
||||
|
||||
use std::net::IpAddr;
|
||||
|
||||
use cidr::{Ipv4Inet, Ipv6Inet};
|
||||
use easytier_proto::common::{FlagsInConfig, SecureModeConfig};
|
||||
use hmac::Hmac;
|
||||
use sha2::Sha256;
|
||||
|
||||
use crate::{
|
||||
config::peers::PeerRuntimeConfig,
|
||||
config::{CoreConfig, IpPrefix, NodeConfig, PeerPolicyConfig, RouteConfig, TrafficConfig},
|
||||
peers::context::{NetworkIdentity, PeerContext, secret_proof_from_secret},
|
||||
};
|
||||
|
||||
pub(crate) trait PeerContextTestExt: PeerContext {
|
||||
fn runtime_config(&self) -> PeerRuntimeConfig {
|
||||
let network_identity = self.network_identity();
|
||||
let hostname = self.hostname();
|
||||
PeerRuntimeConfig {
|
||||
core: CoreConfig {
|
||||
node: NodeConfig {
|
||||
peer_id: None,
|
||||
instance_id: Some(*self.instance_id().as_bytes()),
|
||||
hostname: (!hostname.is_empty()).then_some(hostname),
|
||||
network_name: network_identity.network_name.clone(),
|
||||
},
|
||||
routes: RouteConfig {
|
||||
ipv4: self.ipv4().map(ipv4_inet_to_config),
|
||||
ipv6: self.ipv6().map(ipv6_inet_to_config),
|
||||
..Default::default()
|
||||
},
|
||||
peer_policy: PeerPolicyConfig::default(),
|
||||
traffic: TrafficConfig::default(),
|
||||
},
|
||||
network_identity,
|
||||
stun_info: self.stun_info(),
|
||||
feature_flags: self.feature_flags(),
|
||||
secure_mode: self.secure_mode(),
|
||||
host_routing: self.host_routing_policy(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn ipv4_inet_to_config(value: Ipv4Inet) -> IpPrefix {
|
||||
IpPrefix::new(IpAddr::V4(value.address()), value.network_length())
|
||||
.expect("Ipv4Inet should always have a valid IPv4 prefix length")
|
||||
}
|
||||
|
||||
fn ipv6_inet_to_config(value: Ipv6Inet) -> IpPrefix {
|
||||
IpPrefix::new(IpAddr::V6(value.address()), value.network_length())
|
||||
.expect("Ipv6Inet should always have a valid IPv6 prefix length")
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct NoopPeerContext {
|
||||
network_identity: NetworkIdentity,
|
||||
flags: FlagsInConfig,
|
||||
secure_mode: Option<SecureModeConfig>,
|
||||
}
|
||||
|
||||
impl NoopPeerContext {
|
||||
pub(crate) fn new(network_identity: NetworkIdentity) -> Self {
|
||||
Self {
|
||||
network_identity,
|
||||
flags: FlagsInConfig::default(),
|
||||
secure_mode: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for NoopPeerContext {
|
||||
fn default() -> Self {
|
||||
Self::new(NetworkIdentity::default())
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerContext for NoopPeerContext {
|
||||
fn network_identity(&self) -> NetworkIdentity {
|
||||
self.network_identity.clone()
|
||||
}
|
||||
|
||||
fn flags(&self) -> FlagsInConfig {
|
||||
self.flags.clone()
|
||||
}
|
||||
|
||||
fn secure_mode(&self) -> Option<SecureModeConfig> {
|
||||
self.secure_mode.clone()
|
||||
}
|
||||
|
||||
fn secret_proof(&self, challenge: &[u8]) -> Option<Hmac<Sha256>> {
|
||||
let secret = self.network_identity.network_secret.as_ref()?;
|
||||
secret_proof_from_secret(secret, challenge)
|
||||
}
|
||||
}
|
||||
|
||||
impl PeerContextTestExt for NoopPeerContext {}
|
||||
@@ -0,0 +1,113 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::foundation::time::{Duration, timeout};
|
||||
|
||||
use crate::{
|
||||
packet::{PacketType, ZCPacket},
|
||||
peers::{
|
||||
conn::{peer_conn::PeerConn, peer_map::PeerMap, peer_session::PeerSessionStore},
|
||||
context::NetworkIdentity,
|
||||
create_packet_recv_chan,
|
||||
error::Error,
|
||||
test_support::NoopPeerContext,
|
||||
},
|
||||
tunnel::ring::create_ring_tunnel_pair,
|
||||
};
|
||||
|
||||
impl PeerConn {
|
||||
#[tracing::instrument]
|
||||
async fn do_handshake_as_server(&mut self) -> Result<(), Error> {
|
||||
self.do_handshake_as_server_ext(|_, _| Ok(())).await
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_conn_handshake_over_memory_tunnel() {
|
||||
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 = PeerConn::new(1, client_ctx, client_tunnel, peer_session_store.clone());
|
||||
let mut server = PeerConn::new(2, server_ctx, server_tunnel, peer_session_store);
|
||||
|
||||
let (client_ret, server_ret) = tokio::join!(
|
||||
client.do_handshake_as_client(),
|
||||
server.do_handshake_as_server()
|
||||
);
|
||||
|
||||
client_ret.unwrap();
|
||||
server_ret.unwrap();
|
||||
assert_eq!(client.get_peer_id(), 2);
|
||||
assert_eq!(server.get_peer_id(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_conn_handshake_matches_plaintext_secret_identity() {
|
||||
let peer_session_store = Arc::new(PeerSessionStore::new());
|
||||
let (client_tunnel, server_tunnel) = create_ring_tunnel_pair();
|
||||
let client_ctx = Arc::new(NoopPeerContext::new(NetworkIdentity {
|
||||
network_name: "net".to_string(),
|
||||
network_secret: Some("secret".to_string()),
|
||||
network_secret_digest: None,
|
||||
}));
|
||||
let server_ctx = Arc::new(NoopPeerContext::new(NetworkIdentity {
|
||||
network_name: "net".to_string(),
|
||||
network_secret: Some("secret".to_string()),
|
||||
network_secret_digest: None,
|
||||
}));
|
||||
|
||||
let mut client = PeerConn::new(1, client_ctx, client_tunnel, peer_session_store.clone());
|
||||
let mut server = PeerConn::new(2, server_ctx, server_tunnel, peer_session_store);
|
||||
|
||||
let (client_ret, server_ret) = tokio::join!(
|
||||
client.do_handshake_as_client(),
|
||||
server.do_handshake_as_server()
|
||||
);
|
||||
|
||||
client_ret.unwrap();
|
||||
server_ret.unwrap();
|
||||
assert!(client.matches_local_network_secret());
|
||||
assert!(server.matches_local_network_secret());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn peer_map_forwards_packet_over_memory_tunnel() {
|
||||
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(2, server_ctx.clone(), server_tunnel, peer_session_store);
|
||||
|
||||
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();
|
||||
|
||||
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();
|
||||
|
||||
let mut packet = ZCPacket::new_with_payload(b"hello");
|
||||
packet.fill_peer_manager_hdr(1, 2, PacketType::Data as u8);
|
||||
client_map.send_msg_directly(packet, 2).await.unwrap();
|
||||
|
||||
let received = timeout(Duration::from_secs(1), server_rx.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(received.payload(), b"hello");
|
||||
}
|
||||
@@ -0,0 +1,494 @@
|
||||
use std::{future::Future, sync::Arc};
|
||||
|
||||
use dashmap::DashMap;
|
||||
use futures::future::BoxFuture;
|
||||
|
||||
use crate::config::PeerId;
|
||||
use crate::foundation::stats::{CounterHandle, LabelSet, LabelType, MetricName, StatsManager};
|
||||
use crate::packet::PacketType;
|
||||
use crate::peers::util::shrink_dashmap;
|
||||
use crate::proto::peer_rpc::RoutePeerInfo;
|
||||
|
||||
pub const UNKNOWN_INSTANCE_ID: &str = "unknown";
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
pub enum InstanceLabelKind {
|
||||
To,
|
||||
From,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct TrafficCounters {
|
||||
bytes: CounterHandle,
|
||||
packets: CounterHandle,
|
||||
}
|
||||
|
||||
impl TrafficCounters {
|
||||
fn add_sample(&self, bytes: u64) {
|
||||
self.bytes.add(bytes);
|
||||
self.packets.inc();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum CachedPeerTrafficCounters {
|
||||
Unknown(TrafficCounters),
|
||||
Resolved(TrafficCounters),
|
||||
}
|
||||
|
||||
impl CachedPeerTrafficCounters {
|
||||
fn counters(&self) -> TrafficCounters {
|
||||
match self {
|
||||
CachedPeerTrafficCounters::Unknown(counters)
|
||||
| CachedPeerTrafficCounters::Resolved(counters) => counters.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_resolved(&self) -> bool {
|
||||
matches!(self, CachedPeerTrafficCounters::Resolved(_))
|
||||
}
|
||||
}
|
||||
|
||||
pub struct LogicalTrafficMetrics {
|
||||
stats_mgr: Arc<StatsManager>,
|
||||
network_name: String,
|
||||
instance_bytes_metric: MetricName,
|
||||
instance_packets_metric: MetricName,
|
||||
label_kind: InstanceLabelKind,
|
||||
total: TrafficCounters,
|
||||
per_peer: DashMap<PeerId, CachedPeerTrafficCounters>,
|
||||
}
|
||||
|
||||
impl LogicalTrafficMetrics {
|
||||
pub fn new(
|
||||
stats_mgr: Arc<StatsManager>,
|
||||
network_name: String,
|
||||
total_bytes_metric: MetricName,
|
||||
total_packets_metric: MetricName,
|
||||
instance_bytes_metric: MetricName,
|
||||
instance_packets_metric: MetricName,
|
||||
label_kind: InstanceLabelKind,
|
||||
) -> Self {
|
||||
let label_set =
|
||||
LabelSet::new().with_label_type(LabelType::NetworkName(network_name.clone()));
|
||||
Self {
|
||||
total: TrafficCounters {
|
||||
bytes: stats_mgr.get_counter(total_bytes_metric, label_set.clone()),
|
||||
packets: stats_mgr.get_counter(total_packets_metric, label_set),
|
||||
},
|
||||
stats_mgr,
|
||||
network_name,
|
||||
instance_bytes_metric,
|
||||
instance_packets_metric,
|
||||
label_kind,
|
||||
per_peer: DashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn record_with_resolver<F, Fut>(&self, peer_id: PeerId, bytes: u64, resolver: F)
|
||||
where
|
||||
F: FnOnce() -> Fut,
|
||||
Fut: Future<Output = Option<String>>,
|
||||
{
|
||||
self.total.add_sample(bytes);
|
||||
|
||||
if let Some(entry) = self.per_peer.get(&peer_id)
|
||||
&& entry.value().is_resolved()
|
||||
{
|
||||
entry.value().counters().add_sample(bytes);
|
||||
return;
|
||||
}
|
||||
|
||||
let resolved_instance_id = resolver().await;
|
||||
let counters = self.get_or_update_peer_counters(peer_id, resolved_instance_id.as_deref());
|
||||
counters.add_sample(bytes);
|
||||
}
|
||||
|
||||
fn get_or_update_peer_counters(
|
||||
&self,
|
||||
peer_id: PeerId,
|
||||
resolved_instance_id: Option<&str>,
|
||||
) -> TrafficCounters {
|
||||
match self.per_peer.entry(peer_id) {
|
||||
dashmap::Entry::Occupied(mut entry) => {
|
||||
if entry.get().is_resolved() || resolved_instance_id.is_none() {
|
||||
return entry.get().counters();
|
||||
}
|
||||
let counters = self.build_peer_counters(resolved_instance_id.unwrap());
|
||||
entry.insert(CachedPeerTrafficCounters::Resolved(counters.clone()));
|
||||
counters
|
||||
}
|
||||
dashmap::Entry::Vacant(entry) => {
|
||||
let counters =
|
||||
self.build_peer_counters(resolved_instance_id.unwrap_or(UNKNOWN_INSTANCE_ID));
|
||||
let cached = if resolved_instance_id.is_some() {
|
||||
CachedPeerTrafficCounters::Resolved(counters.clone())
|
||||
} else {
|
||||
CachedPeerTrafficCounters::Unknown(counters.clone())
|
||||
};
|
||||
entry.insert(cached);
|
||||
counters
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn remove_peer(&self, peer_id: PeerId) {
|
||||
self.per_peer.remove(&peer_id);
|
||||
shrink_dashmap(&self.per_peer, None);
|
||||
}
|
||||
|
||||
pub fn clear_peer_cache(&self) {
|
||||
self.per_peer.clear();
|
||||
shrink_dashmap(&self.per_peer, None);
|
||||
}
|
||||
|
||||
fn contains_peer_cache(&self, peer_id: PeerId) -> bool {
|
||||
self.per_peer.contains_key(&peer_id)
|
||||
}
|
||||
|
||||
fn build_peer_counters(&self, instance_id: &str) -> TrafficCounters {
|
||||
let instance_label = match self.label_kind {
|
||||
InstanceLabelKind::To => LabelType::ToInstanceId(instance_id.to_string()),
|
||||
InstanceLabelKind::From => LabelType::FromInstanceId(instance_id.to_string()),
|
||||
};
|
||||
let label_set = LabelSet::new()
|
||||
.with_label_type(LabelType::NetworkName(self.network_name.clone()))
|
||||
.with_label_type(instance_label);
|
||||
TrafficCounters {
|
||||
bytes: self
|
||||
.stats_mgr
|
||||
.get_counter(self.instance_bytes_metric, label_set.clone()),
|
||||
packets: self
|
||||
.stats_mgr
|
||||
.get_counter(self.instance_packets_metric, label_set),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum TrafficKind {
|
||||
Data,
|
||||
Control,
|
||||
}
|
||||
|
||||
pub fn traffic_kind(packet_type: u8) -> TrafficKind {
|
||||
if packet_type == PacketType::Data as u8
|
||||
|| packet_type == PacketType::KcpSrc as u8
|
||||
|| packet_type == PacketType::KcpDst as u8
|
||||
|| packet_type == PacketType::QuicSrc as u8
|
||||
|| packet_type == PacketType::QuicDst as u8
|
||||
|| packet_type == PacketType::DataWithKcpSrcModified as u8
|
||||
|| packet_type == PacketType::DataWithQuicSrcModified as u8
|
||||
{
|
||||
TrafficKind::Data
|
||||
} else {
|
||||
TrafficKind::Control
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_relay_data_packet_type(packet_type: u8) -> bool {
|
||||
// Relay handshakes are control-plane setup; payload data is blocked by its
|
||||
// original packet type after the session exists.
|
||||
traffic_kind(packet_type) == TrafficKind::Data
|
||||
|| packet_type == PacketType::ForeignNetworkPacket as u8
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct TrafficMetricGroup {
|
||||
data: Arc<LogicalTrafficMetrics>,
|
||||
control: Arc<LogicalTrafficMetrics>,
|
||||
}
|
||||
|
||||
impl TrafficMetricGroup {
|
||||
fn select(&self, kind: TrafficKind) -> &Arc<LogicalTrafficMetrics> {
|
||||
match kind {
|
||||
TrafficKind::Data => &self.data,
|
||||
TrafficKind::Control => &self.control,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type InstanceIdResolver = dyn Fn(PeerId) -> BoxFuture<'static, Option<String>> + Send + Sync;
|
||||
|
||||
pub struct TrafficMetricRecorder {
|
||||
my_peer_id: PeerId,
|
||||
tx_metrics: TrafficMetricGroup,
|
||||
rx_metrics: TrafficMetricGroup,
|
||||
resolve_instance_id: Arc<InstanceIdResolver>,
|
||||
}
|
||||
|
||||
impl TrafficMetricRecorder {
|
||||
pub fn new<F, Fut>(
|
||||
my_peer_id: PeerId,
|
||||
tx_data: Arc<LogicalTrafficMetrics>,
|
||||
tx_control: Arc<LogicalTrafficMetrics>,
|
||||
rx_data: Arc<LogicalTrafficMetrics>,
|
||||
rx_control: Arc<LogicalTrafficMetrics>,
|
||||
resolve_instance_id: F,
|
||||
) -> Self
|
||||
where
|
||||
F: Fn(PeerId) -> Fut + Send + Sync + 'static,
|
||||
Fut: Future<Output = Option<String>> + Send + 'static,
|
||||
{
|
||||
Self {
|
||||
my_peer_id,
|
||||
tx_metrics: TrafficMetricGroup {
|
||||
data: tx_data,
|
||||
control: tx_control,
|
||||
},
|
||||
rx_metrics: TrafficMetricGroup {
|
||||
data: rx_data,
|
||||
control: rx_control,
|
||||
},
|
||||
resolve_instance_id: Arc::new(move |peer_id| Box::pin(resolve_instance_id(peer_id))),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn record_tx(&self, peer_id: PeerId, packet_type: u8, bytes: u64) {
|
||||
if peer_id == self.my_peer_id {
|
||||
return;
|
||||
}
|
||||
self.tx_metrics
|
||||
.select(traffic_kind(packet_type))
|
||||
.record_with_resolver(peer_id, bytes, || self.resolve_instance_id(peer_id))
|
||||
.await;
|
||||
}
|
||||
|
||||
pub async fn record_rx(&self, peer_id: PeerId, packet_type: u8, bytes: u64) {
|
||||
if peer_id == self.my_peer_id {
|
||||
return;
|
||||
}
|
||||
self.rx_metrics
|
||||
.select(traffic_kind(packet_type))
|
||||
.record_with_resolver(peer_id, bytes, || self.resolve_instance_id(peer_id))
|
||||
.await;
|
||||
}
|
||||
|
||||
pub fn remove_peer(&self, peer_id: PeerId) {
|
||||
self.tx_metrics.data.remove_peer(peer_id);
|
||||
self.tx_metrics.control.remove_peer(peer_id);
|
||||
self.rx_metrics.data.remove_peer(peer_id);
|
||||
self.rx_metrics.control.remove_peer(peer_id);
|
||||
}
|
||||
|
||||
pub fn clear_peer_cache(&self) {
|
||||
self.tx_metrics.data.clear_peer_cache();
|
||||
self.tx_metrics.control.clear_peer_cache();
|
||||
self.rx_metrics.data.clear_peer_cache();
|
||||
self.rx_metrics.control.clear_peer_cache();
|
||||
}
|
||||
|
||||
pub fn contains_peer_cache(&self, peer_id: PeerId) -> bool {
|
||||
self.tx_metrics.data.contains_peer_cache(peer_id)
|
||||
|| self.tx_metrics.control.contains_peer_cache(peer_id)
|
||||
|| self.rx_metrics.data.contains_peer_cache(peer_id)
|
||||
|| self.rx_metrics.control.contains_peer_cache(peer_id)
|
||||
}
|
||||
|
||||
fn resolve_instance_id(&self, peer_id: PeerId) -> BoxFuture<'static, Option<String>> {
|
||||
(self.resolve_instance_id)(peer_id)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn route_peer_info_instance_id(route_peer_info: &RoutePeerInfo) -> Option<String> {
|
||||
let instance_id = route_peer_info.inst_id.as_ref()?;
|
||||
let instance_id: uuid::Uuid = (*instance_id).into();
|
||||
if instance_id.is_nil() {
|
||||
None
|
||||
} else {
|
||||
Some(instance_id.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
impl LogicalTrafficMetrics {
|
||||
fn peer_cache_size(&self) -> usize {
|
||||
self.per_peer.len()
|
||||
}
|
||||
}
|
||||
|
||||
fn network_labels(network_name: &str) -> LabelSet {
|
||||
LabelSet::new().with_label_type(LabelType::NetworkName(network_name.to_string()))
|
||||
}
|
||||
|
||||
fn to_instance_labels(network_name: &str, instance_id: &str) -> LabelSet {
|
||||
LabelSet::new()
|
||||
.with_label_type(LabelType::NetworkName(network_name.to_string()))
|
||||
.with_label_type(LabelType::ToInstanceId(instance_id.to_string()))
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn logical_traffic_metrics_upgrade_unknown_instance_label() {
|
||||
let stats_mgr = Arc::new(StatsManager::new());
|
||||
let metrics = LogicalTrafficMetrics::new(
|
||||
stats_mgr.clone(),
|
||||
"default".to_string(),
|
||||
MetricName::TrafficBytesTx,
|
||||
MetricName::TrafficPacketsTx,
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
MetricName::TrafficPacketsTxByInstance,
|
||||
InstanceLabelKind::To,
|
||||
);
|
||||
let peer_id = 42;
|
||||
let resolved_instance_id = "87ede5a2-9c3d-492d-9bbe-989b9d07e742";
|
||||
|
||||
metrics
|
||||
.record_with_resolver(peer_id, 100, || async { None })
|
||||
.await;
|
||||
metrics
|
||||
.record_with_resolver(peer_id, 200, || async {
|
||||
Some(resolved_instance_id.to_string())
|
||||
})
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(MetricName::TrafficBytesTx, &network_labels("default"))
|
||||
.unwrap()
|
||||
.value,
|
||||
300
|
||||
);
|
||||
assert!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficBytesTx,
|
||||
&to_instance_labels("default", UNKNOWN_INSTANCE_ID),
|
||||
)
|
||||
.is_none()
|
||||
);
|
||||
assert!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficBytesTx,
|
||||
&to_instance_labels("default", resolved_instance_id),
|
||||
)
|
||||
.is_none()
|
||||
);
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
&to_instance_labels("default", UNKNOWN_INSTANCE_ID),
|
||||
)
|
||||
.unwrap()
|
||||
.value,
|
||||
100
|
||||
);
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
&to_instance_labels("default", resolved_instance_id),
|
||||
)
|
||||
.unwrap()
|
||||
.value,
|
||||
200
|
||||
);
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficPacketsTxByInstance,
|
||||
&to_instance_labels("default", UNKNOWN_INSTANCE_ID),
|
||||
)
|
||||
.unwrap()
|
||||
.value,
|
||||
1
|
||||
);
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficPacketsTxByInstance,
|
||||
&to_instance_labels("default", resolved_instance_id),
|
||||
)
|
||||
.unwrap()
|
||||
.value,
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn logical_traffic_metrics_remove_peer_clears_cached_counters() {
|
||||
let stats_mgr = Arc::new(StatsManager::new());
|
||||
let metrics = LogicalTrafficMetrics::new(
|
||||
stats_mgr.clone(),
|
||||
"default".to_string(),
|
||||
MetricName::TrafficBytesTx,
|
||||
MetricName::TrafficPacketsTx,
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
MetricName::TrafficPacketsTxByInstance,
|
||||
InstanceLabelKind::To,
|
||||
);
|
||||
let peer_id = 42;
|
||||
let resolved_instance_id = "87ede5a2-9c3d-492d-9bbe-989b9d07e742";
|
||||
|
||||
metrics
|
||||
.record_with_resolver(peer_id, 100, || async {
|
||||
Some(resolved_instance_id.to_string())
|
||||
})
|
||||
.await;
|
||||
metrics.remove_peer(peer_id);
|
||||
metrics
|
||||
.record_with_resolver(peer_id, 200, || async { None })
|
||||
.await;
|
||||
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(MetricName::TrafficBytesTx, &network_labels("default"))
|
||||
.unwrap()
|
||||
.value,
|
||||
300
|
||||
);
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
&to_instance_labels("default", resolved_instance_id),
|
||||
)
|
||||
.unwrap()
|
||||
.value,
|
||||
100
|
||||
);
|
||||
assert_eq!(
|
||||
stats_mgr
|
||||
.get_metric(
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
&to_instance_labels("default", UNKNOWN_INSTANCE_ID),
|
||||
)
|
||||
.unwrap()
|
||||
.value,
|
||||
200
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn logical_traffic_metrics_clear_peer_cache_resets_all_cached_peers() {
|
||||
let stats_mgr = Arc::new(StatsManager::new());
|
||||
let metrics = LogicalTrafficMetrics::new(
|
||||
stats_mgr,
|
||||
"default".to_string(),
|
||||
MetricName::TrafficBytesTx,
|
||||
MetricName::TrafficPacketsTx,
|
||||
MetricName::TrafficBytesTxByInstance,
|
||||
MetricName::TrafficPacketsTxByInstance,
|
||||
InstanceLabelKind::To,
|
||||
);
|
||||
|
||||
metrics
|
||||
.record_with_resolver(1, 100, || async {
|
||||
Some("87ede5a2-9c3d-492d-9bbe-989b9d07e742".to_string())
|
||||
})
|
||||
.await;
|
||||
metrics
|
||||
.record_with_resolver(2, 200, || async { None })
|
||||
.await;
|
||||
|
||||
assert_eq!(metrics.peer_cache_size(), 2);
|
||||
|
||||
metrics.clear_peer_cache();
|
||||
|
||||
assert_eq!(metrics.peer_cache_size(), 0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
use std::hash::Hash;
|
||||
|
||||
use dashmap::DashMap;
|
||||
|
||||
pub(crate) fn shrink_dashmap<K: Eq + Hash, V>(map: &DashMap<K, V>, threshold: Option<usize>) {
|
||||
let threshold = threshold.unwrap_or(16);
|
||||
if map.capacity() - map.len() > threshold {
|
||||
map.shrink_to_fit();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
//! Relay-network whitelist matching shared by the peer context and the
|
||||
//! foreign-network manager. Kept in the peers kernel so both can depend on it
|
||||
//! without depending on each other.
|
||||
|
||||
pub(crate) fn check_network_in_relay_whitelist(
|
||||
relay_network_whitelist: &str,
|
||||
network_name: &str,
|
||||
) -> Result<(), anyhow::Error> {
|
||||
if relay_network_whitelist
|
||||
.split(' ')
|
||||
.map(wildmatch::WildMatch::new)
|
||||
.any(|whitelist| whitelist.matches(network_name))
|
||||
{
|
||||
Ok(())
|
||||
} else {
|
||||
Err(anyhow::anyhow!("network {} not in whitelist", network_name))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user