Fix credential ospf logic, fix udp subnet proxy loop protection (#2315)

This commit is contained in:
KKRainbow
2026-06-07 12:40:09 +08:00
committed by GitHub
parent 793b57c2a1
commit e38b1354b3
24 changed files with 1892 additions and 703 deletions
+8 -8
View File
@@ -278,20 +278,20 @@ pub struct NetworkIdentity {
pub enum ConfigSource {
#[default]
User,
Webhook,
Web,
}
impl ConfigSource {
pub fn as_str(self) -> &'static str {
match self {
Self::User => "user",
Self::Webhook => "webhook",
Self::Web => "web",
}
}
pub fn from_rpc(source: i32) -> Option<Self> {
match RpcConfigSource::try_from(source).ok() {
Some(RpcConfigSource::Webhook) => Some(Self::Webhook),
Some(RpcConfigSource::Web) => Some(Self::Web),
Some(RpcConfigSource::User) => Some(Self::User),
_ => None,
}
@@ -300,7 +300,7 @@ impl ConfigSource {
pub fn to_rpc(self) -> i32 {
match self {
Self::User => RpcConfigSource::User as i32,
Self::Webhook => RpcConfigSource::Webhook as i32,
Self::Web => RpcConfigSource::Web as i32,
}
}
}
@@ -311,7 +311,7 @@ impl std::str::FromStr for ConfigSource {
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s {
"user" => Ok(Self::User),
"webhook" => Ok(Self::Webhook),
"web" => Ok(Self::Web),
other => Err(format!("unknown network config source: {other}")),
}
}
@@ -1357,14 +1357,14 @@ stun_servers = [
let config = TomlConfigLoader::default();
assert_eq!(config.get_network_config_source(), ConfigSource::User);
config.set_network_config_source(Some(ConfigSource::Webhook));
config.set_network_config_source(Some(ConfigSource::Web));
let dumped = config.dump();
assert!(dumped.contains("[source]"));
assert!(dumped.contains("source = \"webhook\""));
assert!(dumped.contains("source = \"web\""));
let loaded = TomlConfigLoader::new_from_str(&dumped).unwrap();
assert_eq!(loaded.get_network_config_source(), ConfigSource::Webhook);
assert_eq!(loaded.get_network_config_source(), ConfigSource::Web);
}
#[test]
+334 -7
View File
@@ -40,6 +40,16 @@ use crate::{
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
struct UdpNatKey {
src_socket: SocketAddr,
dst_socket: SocketAddr,
}
impl UdpNatKey {
fn new(src_socket: SocketAddr, dst_socket: SocketAddr) -> Self {
Self {
src_socket,
dst_socket,
}
}
}
#[derive(Debug)]
@@ -204,13 +214,23 @@ impl UdpNatEntry {
self_clone.mark_active();
if src_v4.ip().is_loopback() {
src_v4.set_ip(virtual_ipv4);
let has_mapped_dst = real_ipv4 != mapped_ipv4;
let mut reply_src_ip = *src_v4.ip();
// Preserve the existing priority for proxy rules that expose a
// real loopback address as a mapped address. Other loopback
// replies come from local delivery to 127.0.0.1 for the local
// virtual IP and may need the mapped rewrite below.
if has_mapped_dst && reply_src_ip == real_ipv4 {
reply_src_ip = mapped_ipv4;
} else if reply_src_ip.is_loopback() {
reply_src_ip = virtual_ipv4;
}
if *src_v4.ip() == real_ipv4 {
src_v4.set_ip(mapped_ipv4);
if has_mapped_dst && reply_src_ip == real_ipv4 {
reply_src_ip = mapped_ipv4;
}
src_v4.set_ip(reply_src_ip);
let Ok(_) = Self::compose_ipv4_packet(
&self_clone,
@@ -321,9 +341,10 @@ impl UdpProxy {
"udp nat packet request received"
);
let nat_key = UdpNatKey {
src_socket: SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()),
};
let nat_key = UdpNatKey::new(
SocketAddr::new(ipv4.get_source().into(), udp_packet.get_source()),
SocketAddr::new(ipv4.get_destination().into(), udp_packet.get_destination()),
);
let nat_entry = self
.nat_table
.entry(nat_key)
@@ -487,3 +508,309 @@ impl Drop for UdpProxy {
}
}
}
#[cfg(test)]
mod tests {
use std::{
net::{Ipv4Addr, SocketAddr},
sync::Arc,
time::Duration,
};
use pnet::packet::{
MutablePacket, Packet,
ip::IpNextHeaderProtocols,
ipv4::{self, Ipv4Packet, MutableIpv4Packet},
udp::{self, MutableUdpPacket, UdpPacket},
};
use tokio::{net::UdpSocket, sync::mpsc::Receiver, time::timeout};
use crate::{
common::{config::ConfigLoader, global_ctx::tests::get_mock_global_ctx},
peers::{
create_packet_recv_chan,
peer_manager::{PeerManager, RouteAlgoType},
},
tunnel::packet_def::{PacketType, ZCPacket},
};
use super::UdpProxy;
fn build_udp_proxy_packet(
src_ip: Ipv4Addr,
src_port: u16,
dst_socket: SocketAddr,
payload: &[u8],
) -> ZCPacket {
let SocketAddr::V4(dst_socket) = dst_socket else {
panic!("test only builds IPv4 UDP packets");
};
let dst_ip = *dst_socket.ip();
let mut packet = vec![0; 20 + 8 + payload.len()];
let packet_len = packet.len() as u16;
{
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
ipv4_packet.set_version(4);
ipv4_packet.set_header_length(5);
ipv4_packet.set_total_length(packet_len);
ipv4_packet.set_ttl(64);
ipv4_packet.set_next_level_protocol(IpNextHeaderProtocols::Udp);
ipv4_packet.set_source(src_ip);
ipv4_packet.set_destination(dst_ip);
}
{
let mut udp_packet = MutableUdpPacket::new(&mut packet[20..]).unwrap();
udp_packet.set_source(src_port);
udp_packet.set_destination(dst_socket.port());
udp_packet.set_length((8 + payload.len()) as u16);
udp_packet.payload_mut().copy_from_slice(payload);
udp_packet.set_checksum(udp::ipv4_checksum(
&udp_packet.to_immutable(),
&src_ip,
&dst_ip,
));
}
{
let mut ipv4_packet = MutableIpv4Packet::new(&mut packet).unwrap();
ipv4_packet.set_checksum(ipv4::checksum(&ipv4_packet.to_immutable()));
}
let mut packet = ZCPacket::new_with_payload(&packet);
packet.fill_peer_manager_hdr(1009867077, 3831440917, PacketType::Data as u8);
packet
}
async fn wait_proxy_cidr_loaded(proxy: &UdpProxy) {
timeout(Duration::from_secs(1), async {
while proxy.cidr_set.is_empty() {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.unwrap();
}
async fn recv_payload(socket: &UdpSocket) -> (Vec<u8>, SocketAddr) {
let mut buf = [0; 64];
let (len, addr) = timeout(Duration::from_secs(1), socket.recv_from(&mut buf))
.await
.unwrap()
.unwrap();
(buf[..len].to_vec(), addr)
}
async fn recv_response_packet(receiver: &mut Receiver<ZCPacket>) -> ZCPacket {
timeout(Duration::from_secs(1), receiver.recv())
.await
.unwrap()
.unwrap()
}
fn assert_udp_response(
packet: ZCPacket,
src_socket: SocketAddr,
dst_ip: Ipv4Addr,
dst_port: u16,
payload: &[u8],
) {
let SocketAddr::V4(src_socket) = src_socket else {
panic!("test only checks IPv4 UDP packets");
};
let ipv4_packet = Ipv4Packet::new(packet.payload()).unwrap();
assert_eq!(ipv4_packet.get_source(), *src_socket.ip());
assert_eq!(ipv4_packet.get_destination(), dst_ip);
let udp_packet = UdpPacket::new(ipv4_packet.payload()).unwrap();
assert_eq!(udp_packet.get_source(), src_socket.port());
assert_eq!(udp_packet.get_destination(), dst_port);
assert_eq!(udp_packet.payload(), payload);
}
async fn stop_nat_entries(proxy: &UdpProxy) {
let nat_socket_addrs = proxy
.nat_table
.iter()
.filter_map(|entry| {
entry
.socket
.as_ref()
.and_then(|socket| socket.local_addr().ok())
.map(|addr| SocketAddr::from((Ipv4Addr::LOCALHOST, addr.port())))
})
.collect::<Vec<_>>();
for entry in proxy.nat_table.iter() {
entry.stop();
}
let wake_socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
for addr in nat_socket_addrs {
let _ = wake_socket.send_to(b"wake", addr).await;
}
}
#[tokio::test]
async fn udp_proxy_rewrites_unmapped_loopback_reply_to_virtual_ip() {
let global_ctx = get_mock_global_ctx();
global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap()));
global_ctx
.config
.add_proxy_cidr("127.0.0.1/32".parse().unwrap(), None)
.unwrap();
let (packet_sender, _packet_receiver) = create_packet_recv_chan();
let peer_manager = Arc::new(PeerManager::new(
RouteAlgoType::Ospf,
global_ctx.clone(),
packet_sender,
));
let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap();
wait_proxy_cidr_loaded(&proxy).await;
let mut response_receiver = proxy.receiver.lock().await.take().unwrap();
let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let real_dst_port = real_dst.local_addr().unwrap().port();
let dst_socket = SocketAddr::from((Ipv4Addr::LOCALHOST, real_dst_port));
let src_ip = Ipv4Addr::new(10, 144, 144, 206);
let src_port = 53864;
let packet = build_udp_proxy_packet(src_ip, src_port, dst_socket, b"request");
assert!(proxy.try_handle_packet(&packet).await.is_some());
let (payload, nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"request");
real_dst.send_to(b"reply", nat_socket).await.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
SocketAddr::from((Ipv4Addr::new(10, 144, 144, 204), real_dst_port)),
src_ip,
src_port,
b"reply",
);
stop_nat_entries(&proxy).await;
}
#[tokio::test]
async fn udp_proxy_maps_local_virtual_destination_reply_to_mapped_source() {
let global_ctx = get_mock_global_ctx();
global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap()));
global_ctx
.config
.add_proxy_cidr(
"10.144.144.204/32".parse().unwrap(),
Some("10.10.10.3/32".parse().unwrap()),
)
.unwrap();
let (packet_sender, _packet_receiver) = create_packet_recv_chan();
let peer_manager = Arc::new(PeerManager::new(
RouteAlgoType::Ospf,
global_ctx.clone(),
packet_sender,
));
let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap();
wait_proxy_cidr_loaded(&proxy).await;
let mut response_receiver = proxy.receiver.lock().await.take().unwrap();
let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let real_dst_port = real_dst.local_addr().unwrap().port();
let mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 3), real_dst_port));
let src_ip = Ipv4Addr::new(10, 144, 144, 206);
let src_port = 53864;
let packet = build_udp_proxy_packet(src_ip, src_port, mapped_dst, b"request");
assert!(proxy.try_handle_packet(&packet).await.is_some());
let (payload, nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"request");
real_dst.send_to(b"reply", nat_socket).await.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
mapped_dst,
src_ip,
src_port,
b"reply",
);
stop_nat_entries(&proxy).await;
}
#[tokio::test]
async fn udp_proxy_separates_same_source_port_to_multiple_mapped_destinations() {
let global_ctx = get_mock_global_ctx();
global_ctx.set_ipv4(Some("10.144.144.204/24".parse().unwrap()));
global_ctx
.config
.add_proxy_cidr(
"127.0.0.1/32".parse().unwrap(),
Some("10.10.10.1/32".parse().unwrap()),
)
.unwrap();
global_ctx
.config
.add_proxy_cidr(
"127.0.0.1/32".parse().unwrap(),
Some("10.10.10.2/32".parse().unwrap()),
)
.unwrap();
let (packet_sender, _packet_receiver) = create_packet_recv_chan();
let peer_manager = Arc::new(PeerManager::new(
RouteAlgoType::Ospf,
global_ctx.clone(),
packet_sender,
));
let proxy = UdpProxy::new(global_ctx, peer_manager).unwrap();
wait_proxy_cidr_loaded(&proxy).await;
let mut response_receiver = proxy.receiver.lock().await.take().unwrap();
let real_dst = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
let real_dst_port = real_dst.local_addr().unwrap().port();
let first_mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 1), real_dst_port));
let second_mapped_dst = SocketAddr::from((Ipv4Addr::new(10, 10, 10, 2), real_dst_port));
let src_ip = Ipv4Addr::new(10, 144, 144, 206);
let src_port = 53864;
let first_packet = build_udp_proxy_packet(src_ip, src_port, first_mapped_dst, b"first");
assert!(proxy.try_handle_packet(&first_packet).await.is_some());
let (payload, first_nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"first");
let second_packet = build_udp_proxy_packet(src_ip, src_port, second_mapped_dst, b"second");
assert!(proxy.try_handle_packet(&second_packet).await.is_some());
let (payload, second_nat_socket) = recv_payload(&real_dst).await;
assert_eq!(payload, b"second");
assert_eq!(proxy.nat_table.len(), 2);
real_dst
.send_to(b"first-reply", first_nat_socket)
.await
.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
first_mapped_dst,
src_ip,
src_port,
b"first-reply",
);
real_dst
.send_to(b"second-reply", second_nat_socket)
.await
.unwrap();
assert_udp_response(
recv_response_packet(&mut response_receiver).await,
second_mapped_dst,
src_ip,
src_port,
b"second-reply",
);
stop_nat_entries(&proxy).await;
}
}
+81
View File
@@ -127,6 +127,8 @@ impl CredentialManager {
credential_id: Option<String>,
reusable: bool,
) -> (String, String) {
self.remove_expired_credentials();
let mut credentials = self.credentials.lock().unwrap();
let id = if let Some(id) = credential_id
.map(|x| x.trim().to_string())
@@ -194,6 +196,25 @@ impl CredentialManager {
removed
}
pub fn remove_expired_credentials(&self) -> bool {
self.remove_expired_credentials_at(current_unix_timestamp())
}
fn remove_expired_credentials_at(&self, now: i64) -> bool {
let removed = {
let mut credentials = self.credentials.lock().unwrap();
let before = credentials.len();
credentials.retain(|_, entry| entry.is_active_at(now));
before != credentials.len()
};
if removed {
self.save_to_disk();
}
removed
}
pub fn get_trusted_pubkeys(&self, network_secret: &str) -> Vec<TrustedCredentialPubkeyProof> {
let now = current_unix_timestamp();
@@ -496,6 +517,35 @@ mod tests {
assert_eq!(list.len(), 1);
}
#[test]
fn test_remove_expired_credentials_removes_and_persists() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("creds.json");
let mgr = CredentialManager::new(Some(path.clone()));
mgr.generate_credential_with_id(
vec!["active".to_string()],
false,
vec![],
Duration::from_secs(3600),
Some("active-id".to_string()),
);
mgr.generate_credential_with_id(
vec!["expired".to_string()],
false,
vec![],
Duration::from_secs(0),
Some("expired-id".to_string()),
);
assert!(mgr.remove_expired_credentials());
assert_eq!(mgr.list_credentials().len(), 1);
let reloaded = CredentialManager::new(Some(path));
let list = reloaded.list_credentials();
assert_eq!(list.len(), 1);
assert_eq!(list[0].credential_id, "active-id");
}
#[test]
fn test_generate_with_specified_id_reuses_existing_result() {
let mgr = CredentialManager::new(None);
@@ -528,6 +578,37 @@ mod tests {
assert_eq!(list[0].reusable, Some(true));
}
#[test]
fn test_generate_with_specified_id_replaces_expired_existing_result() {
let mgr = CredentialManager::new(None);
let fixed_id = "fixed-credential-id".to_string();
let (id1, secret1) = mgr.generate_credential_with_id(
vec!["expired".to_string()],
false,
vec![],
Duration::from_secs(0),
Some(fixed_id.clone()),
);
let (id2, secret2) = mgr.generate_credential_with_id(
vec!["fresh".to_string()],
true,
vec!["10.0.0.0/24".to_string()],
Duration::from_secs(3600),
Some(fixed_id.clone()),
);
assert_eq!(id1, fixed_id);
assert_eq!(id2, fixed_id);
assert_ne!(secret1, secret2);
let list = mgr.list_credentials();
assert_eq!(list.len(), 1);
assert_eq!(list[0].credential_id, fixed_id);
assert_eq!(list[0].groups, vec!["fresh".to_string()]);
assert!(list[0].allow_relay);
assert_eq!(list[0].allowed_proxy_cidrs, vec!["10.0.0.0/24".to_string()]);
}
#[test]
fn test_generate_non_reusable_credential() {
let mgr = CredentialManager::new(None);
+136 -1
View File
@@ -503,6 +503,33 @@ impl PeerManager {
});
}
async fn close_untrusted_credential_peers(peer_map: &Arc<PeerMap>, global_ctx: &ArcGlobalCtx) {
let network_name = global_ctx.get_network_name();
for peer_id in peer_map.list_peers() {
if !matches!(
peer_map.get_peer_identity_type(peer_id),
Some(PeerIdentityType::Credential)
) {
continue;
}
let Some(peer) = peer_map.get_peer_by_id(peer_id) else {
continue;
};
let Some(pubkey) = peer.get_peer_public_key() else {
continue;
};
if global_ctx.is_pubkey_trusted(&pubkey, &network_name) {
continue;
}
tracing::warn!(?peer_id, "closing untrusted credential peer");
if let Err(e) = peer_map.close_peer(peer_id).await {
tracing::warn!(?e, ?peer_id, "failed to close untrusted credential peer");
}
}
}
fn build_foreign_network_manager_accessor(
peer_map: &Arc<PeerMap>,
) -> Box<dyn GlobalForeignNetworkAccessor> {
@@ -1849,6 +1876,26 @@ impl PeerManager {
});
}
async fn run_credential_gc_routine(&self) {
let global_ctx = self.global_ctx.clone();
let peer_map = self.peers.clone();
self.tasks.lock().await.spawn(async move {
loop {
if global_ctx.get_network_identity().network_secret.is_some() {
if global_ctx
.get_credential_manager()
.remove_expired_credentials()
{
global_ctx.issue_event(GlobalCtxEvent::CredentialChanged);
}
Self::close_untrusted_credential_peers(&peer_map, &global_ctx).await;
}
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
}
});
}
async fn run_traffic_metrics_gc_routine(&self) {
let mut event_receiver = self.global_ctx.subscribe();
let traffic_metrics = self.traffic_metrics.clone();
@@ -1897,6 +1944,7 @@ impl PeerManager {
self.run_relay_session_gc_routine().await;
self.run_recent_traffic_gc_routine().await;
self.run_peer_session_gc_routine().await;
self.run_credential_gc_routine().await;
self.run_traffic_metrics_gc_routine().await;
self.run_foriegn_network().await;
@@ -2135,6 +2183,7 @@ impl PeerManager {
#[cfg(test)]
mod tests {
use base64::Engine;
use std::{
fmt::Debug,
sync::Arc,
@@ -2164,7 +2213,7 @@ mod tests {
},
},
proto::{
common::{CompressionAlgoPb, NatType},
common::{CompressionAlgoPb, NatType, SecureModeConfig},
peer_rpc::SecureAuthLevel,
},
tunnel::{
@@ -3406,6 +3455,92 @@ mod tests {
// a is client, b is server
}
#[tokio::test]
async fn expired_credential_peer_conn_is_closed_without_ospf() {
let (admin_ch, _admin_rx) = create_packet_recv_chan();
let admin_ctx = get_mock_global_ctx();
admin_ctx.config.set_network_identity(NetworkIdentity::new(
"net1".to_string(),
"secret".to_string(),
));
set_secure_mode_cfg(&admin_ctx, true);
let admin = Arc::new(PeerManager::new(
RouteAlgoType::None,
admin_ctx.clone(),
admin_ch,
));
admin.run().await.unwrap();
let (_cred_id, cred_secret) = admin_ctx.get_credential_manager().generate_credential(
vec![],
false,
vec![],
Duration::from_secs(1),
);
let privkey_bytes: [u8; 32] = base64::engine::general_purpose::STANDARD
.decode(&cred_secret)
.unwrap()
.try_into()
.unwrap();
let private = x25519_dalek::StaticSecret::from(privkey_bytes);
let public = x25519_dalek::PublicKey::from(&private);
let (credential_ch, _credential_rx) = create_packet_recv_chan();
let credential_ctx = get_mock_global_ctx();
credential_ctx
.config
.set_network_identity(NetworkIdentity::new_credential("net1".to_string()));
credential_ctx
.config
.set_secure_mode(Some(SecureModeConfig {
enabled: true,
local_private_key: Some(
base64::engine::general_purpose::STANDARD.encode(private.as_bytes()),
),
local_public_key: Some(
base64::engine::general_purpose::STANDARD.encode(public.as_bytes()),
),
}));
let credential = Arc::new(PeerManager::new(
RouteAlgoType::None,
credential_ctx,
credential_ch,
));
credential.run().await.unwrap();
let credential_peer_id = credential.my_peer_id();
connect_peer_manager(credential.clone(), admin.clone()).await;
wait_for_condition(
|| {
let admin = admin.clone();
async move {
admin
.get_peer_map()
.list_peer_conns(credential_peer_id)
.await
.is_some_and(|conns| !conns.is_empty())
}
},
Duration::from_secs(5),
)
.await;
wait_for_condition(
|| {
let admin = admin.clone();
async move {
admin
.get_peer_map()
.list_peer_conns(credential_peer_id)
.await
.is_none_or(|conns| conns.is_empty())
}
},
Duration::from_secs(5),
)
.await;
}
#[tokio::test]
async fn close_conn_in_foreign_network_client() {
let peer_mgr_server = create_mock_peer_manager_with_name("server".to_string()).await;
+525 -84
View File
@@ -441,6 +441,9 @@ struct SyncedRouteInfo {
// Tracks the currently accepted peer for non-reusable credentials.
// Maps credential pubkey bytes -> peer_id.
non_reusable_credential_owners: DashMap<Vec<u8>, PeerId>,
// Duplicate non-reusable credential peers are kept for OSPF sync and topology
// reachability, but excluded from forwarding until owner election selects them.
suppressed_non_reusable_credential_peers: DashMap<PeerId, ()>,
version: AtomicVersion,
}
@@ -660,6 +663,36 @@ impl SyncedRouteInfo {
}
}
fn replace_suppressed_non_reusable_credential_peers(
&self,
suppressed_peers: BTreeSet<PeerId>,
) -> bool {
let current: BTreeSet<_> = self
.suppressed_non_reusable_credential_peers
.iter()
.map(|entry| *entry.key())
.collect();
if current == suppressed_peers {
return false;
}
self.suppressed_non_reusable_credential_peers
.retain(|peer_id, _| suppressed_peers.contains(peer_id));
for peer_id in suppressed_peers {
self.suppressed_non_reusable_credential_peers
.insert(peer_id, ());
}
self.version.inc();
true
}
fn is_route_suppressed(&self, peer_id: PeerId) -> bool {
self.suppressed_non_reusable_credential_peers
.contains_key(&peer_id)
}
fn update_credential_groups(
&self,
peer_infos: &OrderedHashMap<PeerId, RoutePeerInfo>,
@@ -1233,11 +1266,13 @@ impl SyncedRouteInfo {
where
F: FnMut(PeerId) -> bool,
{
self.verify_and_update_credential_trusts_with_active_peers_protecting(
network_secret,
is_peer_active,
None,
)
let (untrusted_peers, global_trusted_keys, _) = self
.verify_and_update_credential_trusts_with_active_peers_protecting(
network_secret,
is_peer_active,
None,
);
(untrusted_peers, global_trusted_keys)
}
fn verify_and_update_credential_trusts_with_active_peers_protecting<F>(
@@ -1248,6 +1283,7 @@ impl SyncedRouteInfo {
) -> (
Vec<PeerId>,
HashMap<Vec<u8>, crate::common::global_ctx::TrustedKeyMetadata>,
bool,
)
where
F: FnMut(PeerId) -> bool,
@@ -1261,14 +1297,18 @@ impl SyncedRouteInfo {
let (all_trusted, global_trusted_keys) =
self.collect_trusted_credentials(&peer_infos, network_secret, now);
let prev_trusted = self.replace_trusted_credential_pubkeys(&all_trusted);
let (active_non_reusable_owners, duplicate_untrusted_peers) =
let (active_non_reusable_owners, mut duplicate_untrusted_peers) =
self.collect_non_reusable_credential_owners(&peer_infos, &all_trusted, is_peer_active);
if let Some(protected_peer_id) = protected_peer_id {
duplicate_untrusted_peers.remove(&protected_peer_id);
}
self.replace_non_reusable_credential_owners(active_non_reusable_owners);
let suppressed_changed =
self.replace_suppressed_non_reusable_credential_peers(duplicate_untrusted_peers);
self.update_credential_groups(&peer_infos, &all_trusted);
let mut untrusted_peers =
Self::collect_revoked_credential_peers(&peer_infos, &prev_trusted, &all_trusted);
untrusted_peers.extend(duplicate_untrusted_peers);
if let Some(protected_peer_id) = protected_peer_id {
untrusted_peers.remove(&protected_peer_id);
}
@@ -1282,7 +1322,11 @@ impl SyncedRouteInfo {
self.remove_peers(untrusted_peers.iter().copied());
}
(untrusted_peers.into_iter().collect(), global_trusted_keys)
(
untrusted_peers.into_iter().collect(),
global_trusted_keys,
suppressed_changed,
)
}
fn is_admin_peer(&self, info: &RoutePeerInfo) -> bool {
@@ -1327,6 +1371,7 @@ type NextHopMap = DashMap<PeerId, NextHopInfo>;
struct RouteTable {
peer_infos: DashMap<PeerId, RoutePeerInfo>,
next_hop_map: NextHopMap,
suppressed_peer_ids: DashMap<PeerId, ()>,
ipv4_peer_id_map: DashMap<Ipv4Addr, PeerIdVersion>,
ipv6_peer_id_map: DashMap<Ipv6Addr, PeerIdVersion>,
cidr_peer_id_map: ArcSwap<PrefixMap<Ipv4Cidr, PeerIdVersion>>,
@@ -1339,6 +1384,7 @@ impl RouteTable {
RouteTable {
peer_infos: DashMap::new(),
next_hop_map: DashMap::new(),
suppressed_peer_ids: DashMap::new(),
ipv4_peer_id_map: DashMap::new(),
ipv6_peer_id_map: DashMap::new(),
cidr_peer_id_map: ArcSwap::new(Arc::new(PrefixMap::new())),
@@ -1348,6 +1394,13 @@ impl RouteTable {
}
fn get_next_hop(&self, dst_peer_id: PeerId) -> Option<NextHopInfo> {
if self.suppressed_peer_ids.contains_key(&dst_peer_id) {
return None;
}
self.get_topology_next_hop(dst_peer_id)
}
fn get_topology_next_hop(&self, dst_peer_id: PeerId) -> Option<NextHopInfo> {
let cur_version = self.next_hop_map_version.get();
self.next_hop_map.get(&dst_peer_id).and_then(|x| {
if x.version >= cur_version {
@@ -1362,6 +1415,18 @@ impl RouteTable {
self.get_next_hop(peer_id).is_some()
}
fn topology_peer_reachable(&self, peer_id: PeerId) -> bool {
self.get_topology_next_hop(peer_id).is_some()
}
fn sync_suppressed_peer_ids(&self, synced_info: &SyncedRouteInfo) {
self.suppressed_peer_ids
.retain(|peer_id, _| synced_info.is_route_suppressed(*peer_id));
for entry in synced_info.suppressed_non_reusable_credential_peers.iter() {
self.suppressed_peer_ids.insert(*entry.key(), ());
}
}
fn get_udp_nat_type(&self, peer_id: PeerId) -> Option<NatType> {
self.peer_infos
.get(&peer_id)
@@ -1398,21 +1463,24 @@ impl RouteTable {
}
for item in peer_id_to_node_index.iter() {
let src_peer_id = item.key();
let src_peer_id = *item.key();
if src_peer_id != my_peer_id && synced_info.is_route_suppressed(src_peer_id) {
continue;
}
let src_node_idx = item.value();
let connected_peers: BTreeSet<_> = synced_info
.get_connected_peers(*src_peer_id)
.get_connected_peers(src_peer_id)
.unwrap_or_default();
// if avoid relay, just set all outgoing edges to a large value: AVOID_RELAY_COST.
let peer_avoid_relay_data = synced_info.get_avoid_relay_data(*src_peer_id);
let peer_avoid_relay_data = synced_info.get_avoid_relay_data(src_peer_id);
for dst_peer_id in connected_peers.iter() {
let Some(dst_node_idx) = peer_id_to_node_index.get(dst_peer_id) else {
continue;
};
let mut cost = cost_calc.calculate_cost(*src_peer_id, *dst_peer_id) as usize;
let mut cost = cost_calc.calculate_cost(src_peer_id, *dst_peer_id) as usize;
if peer_avoid_relay_data {
cost += AVOID_RELAY_COST;
}
@@ -1431,20 +1499,21 @@ impl RouteTable {
v.version >= cur_version
});
self.peer_infos.retain(|k, _| {
// remove peer info for peers we cannot reach.
self.next_hop_map.contains_key(k)
// remove peer info for peers we cannot forward to.
self.peer_reachable(*k)
});
self.ipv4_peer_id_map.retain(|_, v| {
// remove ipv4 map for peers we cannot reach.
self.next_hop_map.contains_key(&v.peer_id)
// remove ipv4 map for peers we cannot forward to.
self.peer_reachable(v.peer_id)
});
self.ipv6_peer_id_map.retain(|_, v| {
// remove ipv6 map for peers we cannot reach.
self.next_hop_map.contains_key(&v.peer_id)
// remove ipv6 map for peers we cannot forward to.
self.peer_reachable(v.peer_id)
});
shrink_dashmap(&self.peer_infos, None);
shrink_dashmap(&self.next_hop_map, None);
shrink_dashmap(&self.suppressed_peer_ids, None);
shrink_dashmap(&self.ipv4_peer_id_map, None);
shrink_dashmap(&self.ipv6_peer_id_map, None);
}
@@ -1545,6 +1614,7 @@ impl RouteTable {
cost_calc: &T,
) {
let version = synced_info.version.get();
self.sync_suppressed_peer_ids(synced_info);
let local_proxy_cidrs = synced_info
.peer_infos
@@ -1594,6 +1664,10 @@ impl RouteTable {
}
let peer_id = item.key();
if !self.peer_reachable(*peer_id) {
continue;
}
let Some(info) = synced_info.peer_infos.read().get(peer_id).cloned() else {
continue;
};
@@ -1717,6 +1791,7 @@ impl RouteTable {
cidrs_v6 = ?self.cidr_v6_peer_id_map.load(),
"update peer cidr map"
);
self.clean_expired_route_info();
}
fn get_peer_id_for_proxy(&self, ip: &IpAddr) -> Option<PeerId> {
@@ -2147,6 +2222,7 @@ impl PeerRouteServiceImpl {
group_trust_map_cache: DashMap::new(),
trusted_credential_pubkeys: DashMap::new(),
non_reusable_credential_owners: DashMap::new(),
suppressed_non_reusable_credential_peers: DashMap::new(),
version: AtomicVersion::new(),
},
public_ipv6_service: std::sync::Mutex::new(Weak::new()),
@@ -2170,10 +2246,9 @@ impl PeerRouteServiceImpl {
ni.network_secret_digest.map(|d| d.to_vec())
}
#[cfg(test)]
fn is_active_non_reusable_credential_peer(&self, peer_id: PeerId) -> bool {
peer_id == self.my_peer_id
|| self.sessions.contains_key(&peer_id)
|| self.route_table.peer_reachable(peer_id)
peer_id == self.my_peer_id || self.route_table.topology_peer_reachable(peer_id)
}
fn is_credential_node(&self) -> bool {
@@ -2488,7 +2563,7 @@ impl PeerRouteServiceImpl {
};
for item in self.synced_route_info.conn_map.read().iter() {
let src_peer_id = *item.0;
if !self.route_table.peer_reachable(src_peer_id) {
if !self.route_table.topology_peer_reachable(src_peer_id) {
continue;
}
add_to_all_peer_ids(src_peer_id, item.1.version.get());
@@ -2555,7 +2630,7 @@ impl PeerRouteServiceImpl {
}
// do not send unreachable peer info to dst peer.
if !self.route_table.peer_reachable(*peer_id) {
if !self.route_table.topology_peer_reachable(*peer_id) {
unreachable_peers_for_peer_info.insert(*peer_id, peer_info.version);
continue;
}
@@ -2573,7 +2648,7 @@ impl PeerRouteServiceImpl {
return false;
};
if self.route_table.peer_reachable(*peer_id) {
if self.route_table.topology_peer_reachable(*peer_id) {
route_infos.push(peer_info.clone());
}
@@ -2623,7 +2698,7 @@ impl PeerRouteServiceImpl {
continue;
}
if !self.route_table.peer_reachable(*peer_id) {
if !self.route_table.topology_peer_reachable(*peer_id) {
unreachable_peers_for_conn_info.insert(*peer_id, conn_info.version.get());
continue;
}
@@ -2641,7 +2716,7 @@ impl PeerRouteServiceImpl {
return false;
};
if self.route_table.peer_reachable(*peer_id) {
if self.route_table.topology_peer_reachable(*peer_id) {
add_to_conn_peer_list(*peer_id, conn_info);
}
@@ -2699,7 +2774,7 @@ impl PeerRouteServiceImpl {
let my_conn_info_updated = self.update_my_conn_info().await;
let my_foreign_network_updated = self.update_my_foreign_network().await;
let mut untrusted_changed = false;
if my_peer_info_updated {
if my_peer_info_updated || my_conn_info_updated {
untrusted_changed = self.refresh_credential_trusts_and_disconnect().await;
}
@@ -2757,7 +2832,7 @@ impl PeerRouteServiceImpl {
fn refresh_credential_trusts(&self) -> Vec<PeerId> {
let network_identity = self.global_ctx.get_network_identity();
let (untrusted, global_trusted_keys) = self
let (untrusted, global_trusted_keys, _) = self
.synced_route_info
.verify_and_update_credential_trusts_with_active_peers_protecting(
network_identity.network_secret.as_deref(),
@@ -2777,17 +2852,19 @@ impl PeerRouteServiceImpl {
// route table from the latest synced peer/conn state before checking active peers.
self.update_route_table_and_cached_local_conn_bitmap();
let (untrusted, global_trusted_keys) = self
let (untrusted, global_trusted_keys, suppressed_changed) = self
.synced_route_info
.verify_and_update_credential_trusts_with_active_peers_protecting(
network_identity.network_secret.as_deref(),
|peer_id| self.is_active_non_reusable_credential_peer(peer_id),
|peer_id| {
peer_id == self.my_peer_id || self.route_table.topology_peer_reachable(peer_id)
},
Some(self.my_peer_id),
);
self.global_ctx
.update_trusted_keys(global_trusted_keys, &network_identity.network_name);
if !untrusted.is_empty() {
if !untrusted.is_empty() || suppressed_changed {
self.update_route_table_and_cached_local_conn_bitmap();
}
untrusted
@@ -2866,7 +2943,7 @@ impl PeerRouteServiceImpl {
if let Ok(d) = now.duration_since(peer_info.last_update.unwrap().try_into().unwrap())
&& (d > REMOVE_DEAD_PEER_INFO_AFTER
|| (d > REMOVE_UNREACHABLE_PEER_INFO_AFTER
&& !self.route_table.peer_reachable(*peer_id)))
&& !self.route_table.topology_peer_reachable(*peer_id)))
{
to_remove.push(*peer_id);
}
@@ -4173,7 +4250,9 @@ mod tests {
time::{Duration, SystemTime},
};
use super::{NextHopInfo, PeerRoute, REMOVE_DEAD_PEER_INFO_AFTER, RouteConnInfo};
use super::{
NextHopInfo, PeerRoute, REMOVE_DEAD_PEER_INFO_AFTER, RouteConnInfo, SyncRouteSession,
};
use crate::proto::common::TimestampExt;
use crate::{
common::{
@@ -4203,6 +4282,8 @@ mod tests {
},
tunnel::common::tests::wait_for_condition,
};
use base64::Engine as _;
use base64::prelude::BASE64_STANDARD;
struct AuthOnlyInterface {
my_peer_id: PeerId,
@@ -4438,6 +4519,31 @@ mod tests {
peer_info
}
fn make_admin_route_peer_info(
peer_id: PeerId,
credential_key: &[u8],
network_secret: &str,
now: i64,
) -> RoutePeerInfo {
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = peer_id;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.to_vec(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
network_secret,
)];
admin_info
}
fn make_route_conn_info<I>(connected_peers: I, last_update: SystemTime) -> RouteConnInfo
where
I: IntoIterator<Item = PeerId>,
@@ -4887,22 +4993,7 @@ mod tests {
let credential_key = vec![7; 32];
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = 30;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.clone(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
network_secret,
)];
let admin_info = make_admin_route_peer_info(30, &credential_key, network_secret, now);
let mut original_peer = RoutePeerInfo::new();
original_peer.peer_id = 41;
@@ -4953,9 +5044,9 @@ mod tests {
let (second_untrusted, _) = service_impl
.synced_route_info
.verify_and_update_credential_trusts(Some(network_secret));
assert_eq!(second_untrusted, vec![41]);
assert!(second_untrusted.is_empty());
assert!(
!service_impl
service_impl
.synced_route_info
.peer_infos
.read()
@@ -4976,6 +5067,8 @@ mod tests {
.map(|entry| *entry.value()),
Some(39)
);
assert!(service_impl.synced_route_info.is_route_suppressed(41));
assert!(!service_impl.synced_route_info.is_route_suppressed(39));
}
#[tokio::test]
@@ -4991,22 +5084,7 @@ mod tests {
let stale_peer_id = 41;
let replacement_peer_id = 39;
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = 30;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.clone(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
network_secret,
)];
let admin_info = make_admin_route_peer_info(30, &credential_key, network_secret, now);
let mut stale_peer = RoutePeerInfo::new();
stale_peer.peer_id = stale_peer_id;
@@ -5077,6 +5155,292 @@ mod tests {
.map(|entry| *entry.value()),
Some(replacement_peer_id)
);
assert!(
!service_impl
.synced_route_info
.is_route_suppressed(stale_peer_id)
);
assert!(
!service_impl
.synced_route_info
.is_route_suppressed(replacement_peer_id)
);
}
#[tokio::test]
async fn suppressed_non_reusable_credential_peer_stays_synced_and_can_be_reactivated() {
const NETWORK_SECRET: &str = "sec1";
const SELF_PEER_ID: PeerId = 1;
const ADMIN_PEER_ID: PeerId = 30;
const FIRST_PEER_ID: PeerId = 39;
const SECOND_PEER_ID: PeerId = 41;
let service_impl = PeerRouteServiceImpl::new(
SELF_PEER_ID,
get_mock_global_ctx_with_network(Some(NetworkIdentity::new(
"test-net".to_string(),
NETWORK_SECRET.to_string(),
))),
);
let now_unix = SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64;
let now = SystemTime::now();
let credential_key = vec![10; 32];
let mut self_info = RoutePeerInfo::new();
self_info.peer_id = SELF_PEER_ID;
self_info.version = 1;
let admin_info =
make_admin_route_peer_info(ADMIN_PEER_ID, &credential_key, NETWORK_SECRET, now_unix);
let mut first_peer = make_credential_route_peer_info(FIRST_PEER_ID, &credential_key);
first_peer.ipv4_addr = Some(std::net::Ipv4Addr::new(10, 144, 0, 39).into());
let mut second_peer = make_credential_route_peer_info(SECOND_PEER_ID, &credential_key);
second_peer.ipv4_addr = Some(std::net::Ipv4Addr::new(10, 144, 0, 41).into());
second_peer.proxy_cidrs.push("10.244.41.0/24".into());
{
let mut peer_infos = service_impl.synced_route_info.peer_infos.write();
peer_infos.insert(self_info.peer_id, self_info);
peer_infos.insert(admin_info.peer_id, admin_info);
peer_infos.insert(first_peer.peer_id, first_peer);
peer_infos.insert(second_peer.peer_id, second_peer);
}
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(SELF_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(
ADMIN_PEER_ID,
make_route_conn_info([SELF_PEER_ID, FIRST_PEER_ID, SECOND_PEER_ID], now),
);
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
}
service_impl.synced_route_info.version.set(1);
let first_untrusted = service_impl.refresh_credential_trusts_with_current_topology();
assert!(first_untrusted.is_empty());
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(FIRST_PEER_ID)
);
assert!(
service_impl
.synced_route_info
.peer_infos
.read()
.contains_key(&SECOND_PEER_ID)
);
assert!(
service_impl
.synced_route_info
.is_route_suppressed(SECOND_PEER_ID)
);
assert!(
service_impl
.route_table
.topology_peer_reachable(SECOND_PEER_ID)
);
assert!(service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID));
assert!(
service_impl
.route_table
.peer_infos
.contains_key(&FIRST_PEER_ID)
);
assert!(
!service_impl
.route_table
.peer_infos
.contains_key(&SECOND_PEER_ID)
);
assert_eq!(
service_impl
.route_table
.ipv4_peer_id_map
.get(&"10.144.0.41".parse().unwrap())
.map(|entry| entry.peer_id),
None
);
assert_eq!(
service_impl
.route_table
.get_peer_id_for_proxy(&"10.244.41.1".parse().unwrap()),
None
);
let sync_session = SyncRouteSession::new(SELF_PEER_ID, ADMIN_PEER_ID);
let sync_peer_ids: BTreeSet<_> = service_impl
.build_route_info(&sync_session)
.unwrap()
.into_iter()
.map(|info| info.peer_id)
.collect();
assert!(sync_peer_ids.contains(&SECOND_PEER_ID));
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(
ADMIN_PEER_ID,
make_route_conn_info([SELF_PEER_ID, SECOND_PEER_ID], now),
);
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([], now));
conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
}
service_impl.synced_route_info.version.inc();
let second_untrusted = service_impl.refresh_credential_trusts_with_current_topology();
assert!(second_untrusted.is_empty());
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(SECOND_PEER_ID)
);
assert!(
!service_impl
.synced_route_info
.is_route_suppressed(SECOND_PEER_ID)
);
assert!(
service_impl
.route_table
.topology_peer_reachable(SECOND_PEER_ID)
);
assert!(!service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(service_impl.route_table.peer_reachable(SECOND_PEER_ID));
assert!(
!service_impl
.route_table
.peer_infos
.contains_key(&FIRST_PEER_ID)
);
assert!(
service_impl
.route_table
.peer_infos
.contains_key(&SECOND_PEER_ID)
);
assert_eq!(
service_impl
.route_table
.ipv4_peer_id_map
.get(&"10.144.0.41".parse().unwrap())
.map(|entry| entry.peer_id),
Some(SECOND_PEER_ID)
);
assert_eq!(
service_impl
.route_table
.get_peer_id_for_proxy(&"10.244.41.1".parse().unwrap()),
Some(SECOND_PEER_ID)
);
}
#[tokio::test]
async fn suppressed_non_reusable_credential_peer_is_not_transit_next_hop() {
const NETWORK_SECRET: &str = "sec1";
const SELF_PEER_ID: PeerId = 1;
const ADMIN_PEER_ID: PeerId = 30;
const FIRST_PEER_ID: PeerId = 39;
const SECOND_PEER_ID: PeerId = 41;
const DOWNSTREAM_PEER_ID: PeerId = 50;
let service_impl = PeerRouteServiceImpl::new(
SELF_PEER_ID,
get_mock_global_ctx_with_network(Some(NetworkIdentity::new(
"test-net".to_string(),
NETWORK_SECRET.to_string(),
))),
);
let now_unix = SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64;
let now = SystemTime::now();
let credential_key = vec![10; 32];
let mut self_info = RoutePeerInfo::new();
self_info.peer_id = SELF_PEER_ID;
self_info.version = 1;
let admin_info =
make_admin_route_peer_info(ADMIN_PEER_ID, &credential_key, NETWORK_SECRET, now_unix);
let first_peer = make_credential_route_peer_info(FIRST_PEER_ID, &credential_key);
let second_peer = make_credential_route_peer_info(SECOND_PEER_ID, &credential_key);
let mut downstream_peer = RoutePeerInfo::new();
downstream_peer.peer_id = DOWNSTREAM_PEER_ID;
downstream_peer.version = 1;
{
let mut peer_infos = service_impl.synced_route_info.peer_infos.write();
peer_infos.insert(self_info.peer_id, self_info);
peer_infos.insert(admin_info.peer_id, admin_info);
peer_infos.insert(first_peer.peer_id, first_peer);
peer_infos.insert(second_peer.peer_id, second_peer);
peer_infos.insert(downstream_peer.peer_id, downstream_peer);
}
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(SELF_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(
ADMIN_PEER_ID,
make_route_conn_info([SELF_PEER_ID, FIRST_PEER_ID, SECOND_PEER_ID], now),
);
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(
SECOND_PEER_ID,
make_route_conn_info([ADMIN_PEER_ID, DOWNSTREAM_PEER_ID], now),
);
conn_map.insert(
DOWNSTREAM_PEER_ID,
make_route_conn_info([SECOND_PEER_ID], now),
);
}
service_impl.synced_route_info.version.set(1);
let untrusted = service_impl.refresh_credential_trusts_with_current_topology();
assert!(untrusted.is_empty());
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(FIRST_PEER_ID)
);
assert!(
service_impl
.synced_route_info
.is_route_suppressed(SECOND_PEER_ID)
);
assert!(
service_impl
.route_table
.topology_peer_reachable(SECOND_PEER_ID)
);
assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID));
assert!(!service_impl.route_table.peer_reachable(DOWNSTREAM_PEER_ID));
assert!(
service_impl
.route_table
.get_next_hop(DOWNSTREAM_PEER_ID)
.is_none()
);
assert!(
!service_impl
.route_table
.peer_infos
.contains_key(&DOWNSTREAM_PEER_ID)
);
}
#[tokio::test]
@@ -5106,7 +5470,7 @@ mod tests {
},
);
let (untrusted_peers, _) = service_impl
let (untrusted_peers, _, _) = service_impl
.synced_route_info
.verify_and_update_credential_trusts_with_active_peers_protecting(
None,
@@ -5157,22 +5521,8 @@ mod tests {
self_info.peer_id = SELF_PEER_ID;
self_info.version = 1;
let mut admin_info = RoutePeerInfo::new();
admin_info.peer_id = admin_peer_id;
admin_info.version = 1;
admin_info.feature_flag = Some(PeerFeatureFlag {
is_credential_peer: false,
..Default::default()
});
admin_info.trusted_credential_pubkeys = vec![TrustedCredentialPubkeyProof::new_signed(
TrustedCredentialPubkey {
pubkey: credential_key.clone(),
expiry_unix: now + 600,
reusable: Some(false),
..Default::default()
},
NETWORK_SECRET,
)];
let admin_info =
make_admin_route_peer_info(admin_peer_id, &credential_key, NETWORK_SECRET, now);
let stale_peer = make_credential_route_peer_info(stale_peer_id, &credential_key);
let replacement_peer =
@@ -5232,6 +5582,97 @@ mod tests {
);
}
#[tokio::test]
async fn update_my_infos_refreshes_non_reusable_owner_on_conn_change() {
const NETWORK_SECRET: &str = "sec1";
const ADMIN_PEER_ID: PeerId = 30;
const FIRST_PEER_ID: PeerId = 39;
const SECOND_PEER_ID: PeerId = 41;
let global_ctx = get_mock_global_ctx_with_network(Some(NetworkIdentity::new(
"test-net".to_string(),
NETWORK_SECRET.to_string(),
)));
let (_credential_id, credential_secret) = global_ctx
.get_credential_manager()
.generate_credential_with_options(
vec![],
false,
vec![],
Duration::from_secs(3600),
None,
false,
);
let credential_secret_bytes: [u8; 32] = BASE64_STANDARD
.decode(&credential_secret)
.unwrap()
.try_into()
.unwrap();
let credential_secret = x25519_dalek::StaticSecret::from(credential_secret_bytes);
let credential_key = x25519_dalek::PublicKey::from(&credential_secret)
.as_bytes()
.to_vec();
let service_impl = PeerRouteServiceImpl::new(ADMIN_PEER_ID, global_ctx);
let peers = Arc::new(Mutex::new(vec![FIRST_PEER_ID, SECOND_PEER_ID]));
let peer_identity_types = Arc::new(Mutex::new(HashMap::from([
(FIRST_PEER_ID, Some(PeerIdentityType::Credential)),
(SECOND_PEER_ID, Some(PeerIdentityType::Credential)),
])));
*service_impl.interface.lock().await = Some(Box::new(CountingInterface {
my_peer_id: ADMIN_PEER_ID,
peers: peers.clone(),
peer_identity_types,
list_peers_calls: Arc::new(AtomicU32::new(0)),
get_peer_identity_type_calls: Arc::new(AtomicU32::new(0)),
}));
{
let mut peer_infos = service_impl.synced_route_info.peer_infos.write();
peer_infos.insert(
FIRST_PEER_ID,
make_credential_route_peer_info(FIRST_PEER_ID, &credential_key),
);
peer_infos.insert(
SECOND_PEER_ID,
make_credential_route_peer_info(SECOND_PEER_ID, &credential_key),
);
}
let now = SystemTime::now();
{
let mut conn_map = service_impl.synced_route_info.conn_map.write();
conn_map.insert(FIRST_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
conn_map.insert(SECOND_PEER_ID, make_route_conn_info([ADMIN_PEER_ID], now));
}
assert!(service_impl.update_my_infos().await);
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(FIRST_PEER_ID)
);
assert!(service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(!service_impl.route_table.peer_reachable(SECOND_PEER_ID));
*peers.lock() = vec![SECOND_PEER_ID];
service_impl.handle_global_ctx_event(&GlobalCtxEvent::PeerConnRemoved(Default::default()));
assert!(service_impl.update_my_infos().await);
assert_eq!(
service_impl
.synced_route_info
.non_reusable_credential_owners
.get(&credential_key)
.map(|entry| *entry.value()),
Some(SECOND_PEER_ID)
);
assert!(!service_impl.route_table.peer_reachable(FIRST_PEER_ID));
assert!(service_impl.route_table.peer_reachable(SECOND_PEER_ID));
}
#[tokio::test]
async fn sync_route_info_marks_credential_sender_and_filters_entries() {
let peer_mgr = create_mock_pmgr().await;
-3
View File
@@ -1220,9 +1220,6 @@ async fn credential_expiry_disconnects_from_all_admins() {
.await;
tokio::time::sleep(Duration::from_secs(3)).await;
admin_a
.get_global_ctx()
.issue_event(crate::common::global_ctx::GlobalCtxEvent::CredentialChanged);
wait_for_condition(
|| {
+1 -1
View File
@@ -16,7 +16,7 @@ enum NetworkingMethod {
enum ConfigSource {
ConfigSourceUnspecified = 0;
ConfigSourceUser = 1;
ConfigSourceWebhook = 2;
ConfigSourceWeb = 2;
}
message NetworkConfig {
+25 -8
View File
@@ -52,6 +52,17 @@ where
identify: T,
config: NetworkConfig,
save: bool,
) -> Result<(), RemoteClientError<E>> {
self.handle_run_network_instance_with_source(identify, config, save, ConfigSource::User)
.await
}
async fn handle_run_network_instance_with_source(
&self,
identify: T,
config: NetworkConfig,
save: bool,
source: ConfigSource,
) -> Result<(), RemoteClientError<E>> {
let client = self
.get_rpc_client(identify.clone())
@@ -63,7 +74,7 @@ where
inst_id: None,
config: Some(config.clone()),
overwrite: true,
source: ConfigSource::User.to_rpc(),
source: source.to_rpc(),
},
)
.await?;
@@ -74,7 +85,7 @@ where
identify,
resp.inst_id.unwrap_or_default().into(),
config,
ConfigSource::User,
source,
)
.await
.map_err(RemoteClientError::PersistentError)?;
@@ -273,14 +284,20 @@ where
identify: T,
inst_id: uuid::Uuid,
config: NetworkConfig,
) -> Result<(), RemoteClientError<E>> {
self.handle_save_network_config_with_source(identify, inst_id, config, ConfigSource::User)
.await
}
async fn handle_save_network_config_with_source(
&self,
identify: T,
inst_id: uuid::Uuid,
config: NetworkConfig,
source: ConfigSource,
) -> Result<(), RemoteClientError<E>> {
self.get_storage()
.insert_or_update_user_network_config(
identify.clone(),
inst_id,
config,
ConfigSource::User,
)
.insert_or_update_user_network_config(identify.clone(), inst_id, config, source)
.await
.map_err(RemoteClientError::PersistentError)?;
self.get_storage()