diff --git a/vnt-core/src/nat/internal_nat/icmp_nat.rs b/vnt-core/src/nat/internal_nat/icmp_nat.rs index 826895f..49dc76f 100644 --- a/vnt-core/src/nat/internal_nat/icmp_nat.rs +++ b/vnt-core/src/nat/internal_nat/icmp_nat.rs @@ -7,10 +7,18 @@ use pnet_packet::icmp::{IcmpPacket, IcmpTypes}; use pnet_packet::ipv4::Ipv4Packet; use std::collections::HashMap; use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4}; +use std::time::{Duration, Instant}; use tcp_ip::IpStack; use tcp_ip::icmp::IcmpSocket; use tokio::net::UdpSocket; +/// ICMP echo 映射超时:正常 ping 应答在秒级返回,超时条目视为无应答残留 +const ICMP_NAT_TIMEOUT: Duration = Duration::from_secs(60); +const ICMP_NAT_GC_INTERVAL: Duration = Duration::from_secs(60); + +/// (对端地址, identifier, sequence) -> (内网客户端地址, 创建时间) +type IcmpNatMap = HashMap<(Ipv4Addr, Identifier, SequenceNumber), (Ipv4Addr, Instant)>; + pub async fn start_icmp_nat( task_group: &TaskGroup, ip_stack: &IpStack, @@ -49,24 +57,38 @@ async fn task( ) -> anyhow::Result<()> { let mut buf1 = vec![0u8; 65536]; let mut buf2 = vec![0u8; 65536]; - let mut map = HashMap::new(); + let mut map: IcmpNatMap = HashMap::new(); + let mut gc_interval = tokio::time::interval(ICMP_NAT_GC_INTERVAL); loop { tokio::select! { rs = tokio_icmp_socket.recv(&mut buf1) => { let len = rs?; - tokio_icmp_socket_recv(&buf1[..len],&inner_icmp_socket,&map,no_tun,&network).await?; + tokio_icmp_socket_recv(&buf1[..len],&inner_icmp_socket,&mut map,no_tun,&network).await?; } rs = inner_icmp_socket.recv_from_to(&mut buf2) => { let (len,src,dst) = rs?; inner_icmp_socket_recv(&buf2[..len],src,dst,&tokio_icmp_socket,&mut map,no_tun,&network).await?; } + _ = gc_interval.tick() => { + evict_expired(&mut map, Instant::now(), ICMP_NAT_TIMEOUT); + } } } } + +/// 清理超时未收到应答的映射条目,防止 map 无界增长 +fn evict_expired(map: &mut IcmpNatMap, now: Instant, timeout: Duration) { + let before = map.len(); + map.retain(|_, (_, created)| now.duration_since(*created) < timeout); + let evicted = before - map.len(); + if evicted > 0 { + log::debug!("icmp nat evicted {} expired entries", evicted); + } +} async fn tokio_icmp_socket_recv( buf: &[u8], inner_icmp_socket: &IcmpSocket, - map: &HashMap<(Ipv4Addr, Identifier, SequenceNumber), Ipv4Addr>, + map: &mut IcmpNatMap, no_tun: bool, network: &SharedNetworkAddr, ) -> anyhow::Result<()> { @@ -88,7 +110,8 @@ async fn tokio_icmp_socket_recv( let mut src = ipv4.get_source(); let identifier = Identifier::new(u16::from_be_bytes([payload[0], payload[1]])); let sequence_number = SequenceNumber::new(u16::from_be_bytes([payload[2], payload[3]])); - let Some(dst) = map.get(&(src, identifier, sequence_number)) else { + // 收到应答即完成一次 echo 交换,移除映射,避免条目残留 + let Some((dst, _)) = map.remove(&(src, identifier, sequence_number)) else { return Ok(()); }; if no_tun && src == Ipv4Addr::LOCALHOST { @@ -96,7 +119,7 @@ async fn tokio_icmp_socket_recv( } inner_icmp_socket - .send_from_to(ipv4.payload(), src.into(), (*dst).into()) + .send_from_to(ipv4.payload(), src.into(), dst.into()) .await .context("sending ICMPv4 failed")?; Ok(()) @@ -106,7 +129,7 @@ async fn inner_icmp_socket_recv( src: IpAddr, dst: IpAddr, tokio_icmp_socket: &UdpSocket, - map: &mut HashMap<(Ipv4Addr, Identifier, SequenceNumber), Ipv4Addr>, + map: &mut IcmpNatMap, no_tun: bool, network: &SharedNetworkAddr, ) -> anyhow::Result<()> { @@ -131,9 +154,49 @@ async fn inner_icmp_socket_recv( let identifier = Identifier::new(u16::from_be_bytes([payload[0], payload[1]])); let sequence_number = SequenceNumber::new(u16::from_be_bytes([payload[2], payload[3]])); - map.insert((dst, identifier, sequence_number), src); + map.insert((dst, identifier, sequence_number), (src, Instant::now())); tokio_icmp_socket .send_to(buf, SocketAddr::new(dst.into(), 0)) .await?; Ok(()) } + + +#[cfg(test)] +mod tests { + use super::*; + + fn entry( + ip: &str, + id: u16, + seq: u16, + created: Instant, + ) -> ((Ipv4Addr, Identifier, SequenceNumber), (Ipv4Addr, Instant)) { + ( + ( + ip.parse().unwrap(), + Identifier::new(id), + SequenceNumber::new(seq), + ), + (Ipv4Addr::new(10, 0, 0, 1), created), + ) + } + + /// 超时未应答的条目必须被清理,未超时的保留,map 不会无界增长 + #[test] + fn test_evict_expired() { + let now = Instant::now(); + let mut map: IcmpNatMap = HashMap::new(); + let (k_fresh, v_fresh) = entry("8.8.8.8", 1, 1, now); + let (k_stale, v_stale) = + entry("1.1.1.1", 2, 2, now - ICMP_NAT_TIMEOUT - Duration::from_secs(1)); + map.insert(k_fresh, v_fresh); + map.insert(k_stale, v_stale); + + evict_expired(&mut map, now, ICMP_NAT_TIMEOUT); + + assert_eq!(map.len(), 1); + assert!(map.contains_key(&k_fresh)); + assert!(!map.contains_key(&k_stale)); + } +}