diff --git a/easytier/src/connector/direct.rs b/easytier/src/connector/direct.rs index dbde5223..4069130a 100644 --- a/easytier/src/connector/direct.rs +++ b/easytier/src/connector/direct.rs @@ -50,6 +50,7 @@ use url::Host; pub const DIRECT_CONNECTOR_SERVICE_ID: u32 = 1; pub const DIRECT_CONNECTOR_BLACKLIST_TIMEOUT_SEC: u64 = 300; +const MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16; static TESTING: AtomicBool = AtomicBool::new(false); @@ -84,6 +85,56 @@ fn is_usable_public_ipv6_candidate_with_mode( && !ip.is_multicast())) } +fn push_ipv6_hole_punch_candidate( + candidates: &mut Vec, + ip: Ipv6Addr, + global_ctx: &ArcGlobalCtx, + limit: usize, +) { + if candidates.len() >= limit + || !is_usable_public_ipv6_candidate(&ip, global_ctx) + || candidates.contains(&ip) + { + return; + } + candidates.push(ip); +} + +async fn collect_ipv6_hole_punch_candidates(global_ctx: &ArcGlobalCtx) -> Vec { + let mut candidates = Vec::new(); + for ip in global_ctx + .get_stun_info_collector() + .get_stun_info() + .public_ip + .iter() + .filter_map(|ip| ip.parse::().ok()) + { + push_ipv6_hole_punch_candidate( + &mut candidates, + ip, + global_ctx, + MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS, + ); + } + + let ip_list = global_ctx.get_ip_collector().collect_ip_addrs().await; + for ip in ip_list + .interface_ipv6s + .iter() + .chain(ip_list.public_ipv6.iter()) + .map(|ip| Ipv6Addr::from(*ip)) + { + push_ipv6_hole_punch_candidate( + &mut candidates, + ip, + global_ctx, + MAX_IPV6_HOLE_PUNCH_CONNECTOR_ADDRS, + ); + } + + candidates +} + #[async_trait::async_trait] pub trait PeerManagerForDirectConnector { async fn list_peers(&self) -> Vec; @@ -153,7 +204,8 @@ impl DirectConnectorManagerData { async fn remote_send_udp_hole_punch_packet( &self, dst_peer_id: PeerId, - connector_addr: SocketAddr, + connector_addrs: Vec, + preferred_src_ipv6: Option, remote_url: &url::Url, ) -> Result<(), Error> { if !matches_scheme!(remote_url, TunnelScheme::Ip(IpScheme::Udp)) { @@ -184,15 +236,17 @@ impl DirectConnectorManagerData { .send_udp_hole_punch_packet( BaseController::default(), SendUdpHolePunchPacketRequest { + connector_addr: connector_addrs.first().copied().map(Into::into), listener_port: listener_port as u32, - connector_addr: Some(connector_addr.into()), + preferred_src_ipv6: preferred_src_ipv6.map(Into::into), + connector_addrs: connector_addrs.into_iter().map(Into::into).collect(), }, ) .await .with_context(|| { format!( - "do rpc, send udp hole punch packet to peer {} at {}", - dst_peer_id, remote_url + "do rpc, send udp hole punch packet to peer {} at {} with preferred source {:?}", + dst_peer_id, remote_url, preferred_src_ipv6 ) })?; @@ -209,23 +263,41 @@ impl DirectConnectorManagerData { .await .with_context(|| format!("failed to bind local socket for {}", remote_url))?, ); - let connector_ip = self - .global_ctx - .get_stun_info_collector() - .get_stun_info() - .public_ip - .iter() - .filter_map(|ip| ip.parse::().ok()) - .find(|ip| !self.global_ctx.is_ip_easytier_managed_ipv6(ip)); + let connector_ips = collect_ipv6_hole_punch_candidates(&self.global_ctx).await; // ask remote to send v6 hole punch packet // and no matter what the result is, continue to connect - if let Some(connector_ip) = connector_ip { - let connector_addr = - SocketAddr::new(IpAddr::V6(connector_ip), local_socket.local_addr()?.port()); - let _ = self - .remote_send_udp_hole_punch_packet(dst_peer_id, connector_addr, remote_url) - .await; + if !connector_ips.is_empty() { + let local_port = local_socket.local_addr()?.port(); + let connector_addrs = connector_ips + .into_iter() + .map(|ip| SocketAddr::new(IpAddr::V6(ip), local_port)) + .collect::>(); + let preferred_src_ipv6 = match remote_url.host() { + Some(Host::Ipv6(ip)) => Some(ip), + _ => None, + }; + tracing::debug!( + ?connector_addrs, + ?preferred_src_ipv6, + ?remote_url, + "request remote IPv6 hole-punch packets" + ); + if let Err(err) = self + .remote_send_udp_hole_punch_packet( + dst_peer_id, + connector_addrs, + preferred_src_ipv6, + remote_url, + ) + .await + { + tracing::debug!( + ?err, + ?remote_url, + "remote IPv6 hole-punch packet request failed" + ); + } } else { tracing::debug!( ?remote_url, @@ -267,7 +339,7 @@ impl DirectConnectorManagerData { .with_context(|| format!("failed to get udp port mapping for {}", remote_url))?; let _ = self - .remote_send_udp_hole_punch_packet(dst_peer_id, connector_addr, remote_url) + .remote_send_udp_hole_punch_packet(dst_peer_id, vec![connector_addr], None, remote_url) .await; let udp_connector = UdpTunnelConnector::new(remote_url.clone()); @@ -818,7 +890,7 @@ mod tests { tunnel::{IpScheme, TunnelScheme, matches_scheme}, }; - use std::net::{IpAddr, Ipv4Addr, SocketAddr}; + use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use super::{TESTING, mapped_listener_port, resolve_mapped_listener_addrs}; @@ -840,6 +912,27 @@ mod tests { )); } + #[tokio::test] + async fn ipv6_hole_punch_candidates_are_deduped_filtered_and_capped() { + let global_ctx = get_mock_global_ctx(); + let managed_ipv6: cidr::Ipv6Inet = "2001:db8::2/128".parse().unwrap(); + global_ctx.set_public_ipv6_routes(BTreeSet::from([managed_ipv6])); + + let first: Ipv6Addr = "2001:db8::1".parse().unwrap(); + let managed = managed_ipv6.address(); + let second: Ipv6Addr = "2001:db8::3".parse().unwrap(); + let third: Ipv6Addr = "2001:db8::4".parse().unwrap(); + let mut candidates = Vec::new(); + + super::push_ipv6_hole_punch_candidate(&mut candidates, first, &global_ctx, 2); + super::push_ipv6_hole_punch_candidate(&mut candidates, first, &global_ctx, 2); + super::push_ipv6_hole_punch_candidate(&mut candidates, managed, &global_ctx, 2); + super::push_ipv6_hole_punch_candidate(&mut candidates, second, &global_ctx, 2); + super::push_ipv6_hole_punch_candidate(&mut candidates, third, &global_ctx, 2); + + assert_eq!(candidates, vec![first, second]); + } + #[test] fn udp_ipv6_url_matches_hole_punch_branch_condition() { let remote_url: url::Url = "udp://[2001:db8::1]:11010".parse().unwrap(); diff --git a/easytier/src/peers/peer_rpc_service.rs b/easytier/src/peers/peer_rpc_service.rs index 08e82d2d..ec9db867 100644 --- a/easytier/src/peers/peer_rpc_service.rs +++ b/easytier/src/peers/peer_rpc_service.rs @@ -1,7 +1,7 @@ -use std::net::SocketAddr; +use std::net::{IpAddr, Ipv6Addr, SocketAddr}; use crate::{ - common::global_ctx::ArcGlobalCtx, + common::{global_ctx::ArcGlobalCtx, network::IPCollector}, proto::{ common::Void, peer_rpc::{ @@ -12,6 +12,8 @@ use crate::{ tunnel::udp, }; +const MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS: usize = 16; + fn remove_easytier_managed_ipv6s(ret: &mut GetIpListResponse, global_ctx: &ArcGlobalCtx) { ret.interface_ipv6s.retain(|ip| { let ip = std::net::Ipv6Addr::from(*ip); @@ -28,6 +30,86 @@ fn remove_easytier_managed_ipv6s(ret: &mut GetIpListResponse, global_ctx: &ArcGl } } +fn is_usable_preferred_src_ipv6(ip: &Ipv6Addr, global_ctx: &ArcGlobalCtx) -> bool { + !global_ctx.is_ip_easytier_managed_ipv6(ip) + && !ip.is_loopback() + && !ip.is_unspecified() + && !ip.is_unique_local() + && !ip.is_unicast_link_local() + && !ip.is_multicast() +} + +async fn local_preferred_src_ipv6( + global_ctx: &ArcGlobalCtx, + preferred_src_ipv6: Option, +) -> Option { + let preferred_src_ipv6 = preferred_src_ipv6.map(Ipv6Addr::from)?; + if !is_usable_preferred_src_ipv6(&preferred_src_ipv6, global_ctx) { + tracing::debug!( + ?preferred_src_ipv6, + "ignore unusable preferred IPv6 source for udp hole punch" + ); + return None; + } + + let ifaces = IPCollector::collect_interfaces(global_ctx.net_ns.clone(), false).await; + for iface in ifaces { + let is_local = iface.ips.iter().any(|ip| match ip.ip() { + IpAddr::V6(v6) => v6 == preferred_src_ipv6, + IpAddr::V4(_) => false, + }); + if is_local { + tracing::debug!( + ?preferred_src_ipv6, + ifindex = iface.index, + "use preferred IPv6 source for udp hole punch" + ); + return Some(udp::PreferredIpv6Source { + ip: preferred_src_ipv6, + ifindex: iface.index, + }); + } + } + + tracing::debug!( + ?preferred_src_ipv6, + "ignore non-local preferred IPv6 source for udp hole punch" + ); + None +} + +fn connector_addrs_from_request( + req: SendUdpHolePunchPacketRequest, +) -> rpc_types::error::Result<(u16, Vec, Option)> { + let listener_port = u16::try_from(req.listener_port) + .map_err(|_| anyhow::anyhow!("listener_port is out of range: {}", req.listener_port))?; + let mut connector_addrs = req + .connector_addrs + .into_iter() + .map(SocketAddr::from) + .collect::>(); + + if connector_addrs.is_empty() { + connector_addrs.push( + req.connector_addr + .ok_or(anyhow::anyhow!("connector_addr is required"))? + .into(), + ); + } + + let mut deduped = Vec::with_capacity(connector_addrs.len()); + for addr in connector_addrs { + if !deduped.contains(&addr) { + deduped.push(addr); + } + if deduped.len() >= MAX_UDP_HOLE_PUNCH_CONNECTOR_ADDRS { + break; + } + } + + Ok((listener_port, deduped, req.preferred_src_ipv6)) +} + #[derive(Clone)] pub struct DirectConnectorManagerRpcServer { // TODO: this only cache for one src peer, should make it global @@ -67,23 +149,38 @@ impl DirectConnectorRpc for DirectConnectorManagerRpcServer { _: BaseController, req: SendUdpHolePunchPacketRequest, ) -> rpc_types::error::Result { - let listener_port = req.listener_port as u16; - let connector_addr: SocketAddr = req - .connector_addr - .ok_or(anyhow::anyhow!("connector_addr is required"))? - .into(); + let (listener_port, connector_addrs, preferred_src_ipv6) = + connector_addrs_from_request(req)?; + let preferred_src_ipv6 = + local_preferred_src_ipv6(&self.global_ctx, preferred_src_ipv6).await; tracing::info!( - "Sending udp hole punch packet to {} from listener port {}", - connector_addr, - listener_port + ?connector_addrs, + ?preferred_src_ipv6, + listener_port, + "Sending udp hole punch packet" ); // send 3 packets to the connector for _ in 0..3 { - match connector_addr { - SocketAddr::V4(addr) => udp::send_v4_hole_punch_packet(listener_port, addr).await?, - SocketAddr::V6(addr) => udp::send_v6_hole_punch_packet(listener_port, addr).await?, + for connector_addr in &connector_addrs { + let ret = match connector_addr { + SocketAddr::V4(addr) => { + udp::send_v4_hole_punch_packet(listener_port, *addr).await + } + SocketAddr::V6(addr) => { + udp::send_v6_hole_punch_packet(listener_port, *addr, preferred_src_ipv6) + .await + } + }; + if let Err(e) = ret { + tracing::debug!( + ?e, + ?connector_addr, + listener_port, + "send udp hole punch packet failed" + ); + } } tokio::time::sleep(std::time::Duration::from_millis(30)).await; } @@ -99,11 +196,12 @@ impl DirectConnectorManagerRpcServer { #[cfg(test)] mod tests { - use std::collections::BTreeSet; + use std::{collections::BTreeSet, net::SocketAddr}; use crate::{ common::global_ctx::tests::get_mock_global_ctx, - peers::peer_rpc_service::remove_easytier_managed_ipv6s, proto::peer_rpc::GetIpListResponse, + peers::peer_rpc_service::{connector_addrs_from_request, remove_easytier_managed_ipv6s}, + proto::peer_rpc::{GetIpListResponse, SendUdpHolePunchPacketRequest}, }; #[tokio::test] @@ -133,4 +231,58 @@ mod tests { assert_eq!(ip_list.public_ipv6, None); assert_eq!(ip_list.interface_ipv6s, vec![physical_ipv6.into()]); } + + #[test] + fn hole_punch_request_prefers_batch_connector_addrs() { + let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap(); + let first_batch_addr: SocketAddr = "[2001:db8::2]:10002".parse().unwrap(); + let second_batch_addr: SocketAddr = "[2001:db8::3]:10003".parse().unwrap(); + let preferred_src_ipv6: std::net::Ipv6Addr = "2001:db8::4".parse().unwrap(); + + let (listener_port, connector_addrs, preferred_src) = + connector_addrs_from_request(SendUdpHolePunchPacketRequest { + connector_addr: Some(old_addr.into()), + listener_port: 11010, + preferred_src_ipv6: Some(preferred_src_ipv6.into()), + connector_addrs: vec![ + first_batch_addr.into(), + first_batch_addr.into(), + second_batch_addr.into(), + ], + }) + .unwrap(); + + assert_eq!(listener_port, 11010); + assert_eq!(connector_addrs, vec![first_batch_addr, second_batch_addr]); + assert_eq!(preferred_src, Some(preferred_src_ipv6.into())); + } + + #[test] + fn hole_punch_request_falls_back_to_legacy_connector_addr() { + let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap(); + + let (_, connector_addrs, _) = connector_addrs_from_request(SendUdpHolePunchPacketRequest { + connector_addr: Some(old_addr.into()), + listener_port: 11010, + preferred_src_ipv6: None, + connector_addrs: vec![], + }) + .unwrap(); + + assert_eq!(connector_addrs, vec![old_addr]); + } + + #[test] + fn hole_punch_request_rejects_out_of_range_listener_port() { + let old_addr: SocketAddr = "[2001:db8::1]:10001".parse().unwrap(); + + let ret = connector_addrs_from_request(SendUdpHolePunchPacketRequest { + connector_addr: Some(old_addr.into()), + listener_port: u16::MAX as u32 + 1, + preferred_src_ipv6: None, + connector_addrs: vec![], + }); + + assert!(ret.is_err()); + } } diff --git a/easytier/src/proto/peer_rpc.proto b/easytier/src/proto/peer_rpc.proto index 23dc417a..7c8956d9 100644 --- a/easytier/src/proto/peer_rpc.proto +++ b/easytier/src/proto/peer_rpc.proto @@ -189,6 +189,8 @@ message GetIpListResponse { message SendUdpHolePunchPacketRequest { common.SocketAddr connector_addr = 1; uint32 listener_port = 2; + common.Ipv6Addr preferred_src_ipv6 = 3; + repeated common.SocketAddr connector_addrs = 4; } service DirectConnectorRpc { diff --git a/easytier/src/tunnel/mod.rs b/easytier/src/tunnel/mod.rs index 2f0e8032..e3e9de1b 100644 --- a/easytier/src/tunnel/mod.rs +++ b/easytier/src/tunnel/mod.rs @@ -25,6 +25,7 @@ pub mod ring; pub mod stats; pub mod tcp; pub mod udp; +pub(crate) mod udp_src; #[cfg(feature = "faketcp")] pub mod fake_tcp; diff --git a/easytier/src/tunnel/packet_def.rs b/easytier/src/tunnel/packet_def.rs index 10cf5f4d..4ecc64c7 100644 --- a/easytier/src/tunnel/packet_def.rs +++ b/easytier/src/tunnel/packet_def.rs @@ -46,6 +46,8 @@ pub struct V4HolePunchPacket { pub struct V6HolePunchPacket { pub dst_ipv6: [u8; 16], pub dst_port: U16, + pub preferred_src_ipv6: [u8; 16], + pub preferred_src_ifindex: U32, } #[repr(C, packed)] diff --git a/easytier/src/tunnel/udp.rs b/easytier/src/tunnel/udp.rs index 41f5a52d..922e170c 100644 --- a/easytier/src/tunnel/udp.rs +++ b/easytier/src/tunnel/udp.rs @@ -2,13 +2,14 @@ use std::{ fmt::Debug, net::{Ipv4Addr, Ipv6Addr, SocketAddr, SocketAddrV4, SocketAddrV6}, sync::{Arc, Weak}, + time::Duration, }; use anyhow::Context; use async_trait::async_trait; use bytes::BytesMut; use dashmap::DashMap; -use futures::{SinkExt, StreamExt, stream::FuturesUnordered}; +use futures::{StreamExt, stream::FuturesUnordered}; use rand::{Rng, SeedableRng}; use zerocopy::{AsBytes, FromBytes}; @@ -35,6 +36,7 @@ use crate::{ common::{TunnelWrapper, reserve_buf}, packet_def::{UdpPacketType, ZCPacket, ZCPacketType}, ring::RingTunnel, + udp_src, }, }; @@ -43,6 +45,12 @@ pub const UDP_DATA_MTU: usize = 2000; type UdpCloseEventSender = UnboundedSender<(SocketAddr, Option)>; type UdpCloseEventReceiver = UnboundedReceiver<(SocketAddr, Option)>; +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct PreferredIpv6Source { + pub ip: Ipv6Addr, + pub ifindex: u32, +} + fn new_udp_packet(f: F, udp_body: Option<&[u8]>) -> ZCPacket where F: FnOnce(&mut UDPTunnelHeader), @@ -97,11 +105,18 @@ pub fn new_hole_punch_packet(tid: u32, buf_len: u16) -> ZCPacket { ) } -pub fn new_v6_hole_punch_packet(dst: &SocketAddrV6) -> ZCPacket { +pub fn new_v6_hole_punch_packet( + dst: &SocketAddrV6, + preferred_src: Option, +) -> ZCPacket { // generate a 128 bytes vec with random data let mut body = V6HolePunchPacket::default(); body.dst_ipv6.copy_from_slice(&dst.ip().octets()); body.dst_port.set(dst.port()); + if let Some(src) = preferred_src { + body.preferred_src_ipv6.copy_from_slice(&src.ip.octets()); + body.preferred_src_ifindex.set(src.ifindex); + } new_udp_packet( |header| { header.msg_type = UdpPacketType::V6HolePunch as u8; @@ -136,10 +151,18 @@ fn extract_dst_addr_from_v4_hole_punch_packet(buf: &[u8]) -> Option Option { +fn extract_v6_hole_punch_packet(buf: &[u8]) -> Option<(SocketAddrV6, Option)> { let body = V6HolePunchPacket::ref_from_prefix(buf)?; let ip = Ipv6Addr::from(body.dst_ipv6); - Some(SocketAddrV6::new(ip, body.dst_port.get(), 0, 0)) + let preferred_src_ipv6 = Ipv6Addr::from(body.preferred_src_ipv6); + let preferred_src = (!preferred_src_ipv6.is_unspecified()).then_some(PreferredIpv6Source { + ip: preferred_src_ipv6, + ifindex: body.preferred_src_ifindex.get(), + }); + Some(( + SocketAddrV6::new(ip, body.dst_port.get(), 0, 0), + preferred_src, + )) } fn is_stun_packet(b: &[u8]) -> bool { @@ -152,9 +175,10 @@ fn is_stun_packet(b: &[u8]) -> bool { pub async fn send_v6_hole_punch_packet( listener_port: u16, dst_addr: SocketAddrV6, + preferred_src: Option, ) -> Result<(), TunnelError> { let local_socket = UdpSocket::bind("[::1]:0").await?; - let udp_packet = new_v6_hole_punch_packet(&dst_addr); + let udp_packet = new_v6_hole_punch_packet(&dst_addr, preferred_src); let remote_addr = format!("[::1]:{}", listener_port) .parse::() .unwrap(); @@ -369,10 +393,7 @@ impl UdpConnection { } } - pub async fn handle_packet_from_remote( - &mut self, - zc_packet: ZCPacket, - ) -> Result<(), TunnelError> { + pub fn handle_packet_from_remote(&mut self, zc_packet: ZCPacket) -> Result<(), TunnelError> { let header = zc_packet.udp_tunnel_header().unwrap(); let conn_id = header.conn_id.get(); @@ -384,7 +405,14 @@ impl UdpConnection { return Err(TunnelError::ConnIdNotMatch(self.conn_id, conn_id)); } - self.ring_sender.send(zc_packet).await?; + if zc_packet.is_lossy() { + if let Err(e) = self.ring_sender.try_send(zc_packet) { + tracing::trace!(?e, "ring sender full, drop lossy packet"); + } + } else if self.ring_sender.force_send(zc_packet).is_err() { + tracing::trace!("ring sender full, reject non-lossy packet"); + return Err(TunnelError::BufferFull); + } Ok(()) } @@ -431,8 +459,15 @@ impl UdpTunnelListenerData { let socket = self.socket.as_ref().unwrap().clone(); let sack_buf = new_sack_packet(conn_id, magic).into_bytes(); - if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { - tracing::error!(?e, "udp send sack packet error"); + if self + .sock_map + .get(&remote_addr) + .is_some_and(|conn| conn.conn_id == conn_id) + { + if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { + tracing::error!(?e, "udp resend sack packet error"); + } + tracing::debug!(?conn_id, ?remote_addr, "udp duplicate syn, resent sack"); return; } @@ -444,15 +479,43 @@ impl UdpTunnelListenerData { "udp build tunnel for listener" ); - let internal_conn = UdpConnection::new( - socket.clone(), - conn_id, - remote_addr, - RingSink::new(ring_for_recv_udp.clone()), - RingStream::new(ring_for_send_udp.clone()), - self.close_event_sender.clone(), - ); - self.sock_map.insert(remote_addr, internal_conn); + let new_internal_conn = || { + UdpConnection::new( + socket.clone(), + conn_id, + remote_addr, + RingSink::new(ring_for_recv_udp.clone()), + RingStream::new(ring_for_send_udp.clone()), + self.close_event_sender.clone(), + ) + }; + let duplicate_syn = match self.sock_map.entry(remote_addr) { + dashmap::mapref::entry::Entry::Occupied(entry) if entry.get().conn_id == conn_id => { + true + } + dashmap::mapref::entry::Entry::Occupied(mut entry) => { + entry.insert(new_internal_conn()); + false + } + dashmap::mapref::entry::Entry::Vacant(entry) => { + entry.insert(new_internal_conn()); + false + } + }; + if duplicate_syn { + if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { + tracing::error!(?e, "udp resend sack packet error"); + } + tracing::debug!(?conn_id, ?remote_addr, "udp duplicate syn, resent sack"); + return; + } + + if let Err(e) = socket.send_to(&sack_buf, remote_addr).await { + self.sock_map + .remove_if(&remote_addr, |_, conn| conn.conn_id == conn_id); + tracing::error!(?e, "udp send sack packet error"); + return; + } let conn = Box::new(TunnelWrapper::new( Box::new(RingStream::new(ring_for_recv_udp)), @@ -476,7 +539,7 @@ impl UdpTunnelListenerData { } } - async fn do_forward_one_packet_to_conn(&self, zc_packet: ZCPacket, addr: SocketAddr) { + fn do_forward_one_packet_to_conn(&self, zc_packet: ZCPacket, addr: SocketAddr) { let header = zc_packet.udp_tunnel_header().unwrap(); if header.msg_type == UdpPacketType::Syn as u8 { tokio::spawn(Self::handle_new_connect(self.clone(), addr, zc_packet)); @@ -520,23 +583,61 @@ impl UdpTunnelListenerData { tracing::warn!(?addr, "v6 hole punch packet should be sent from ipv6"); return; } - let Some(dst_addr) = extrace_dst_addr_from_hole_punch_packet(zc_packet.udp_payload()) + let Some((dst_addr, preferred_src)) = + extract_v6_hole_punch_packet(zc_packet.udp_payload()) else { tracing::warn!("invalid v6 hole punch packet"); return; }; let socket = self.socket.as_ref().unwrap().clone(); let udp_packet = new_hole_punch_packet(1, 32); - if let Err(e) = socket.try_send_to(&udp_packet.into_bytes(), SocketAddr::V6(dst_addr)) { + let udp_packet = udp_packet.into_bytes(); + let sent_with_src = if let Some(src) = preferred_src { + match udp_src::send_to_with_src_ipv6( + &socket, + src.ip, + src.ifindex, + dst_addr, + &udp_packet, + ) { + Ok(ret) => { + tracing::debug!( + ?src, + ?dst_addr, + ?ret, + "udp forward packet send hole punch packet with preferred ipv6 source" + ); + true + } + Err(e) => { + tracing::debug!( + ?src, + ?dst_addr, + ?e, + "udp forward packet preferred ipv6 source failed, falling back" + ); + false + } + } + } else { + false + }; + if !sent_with_src + && let Err(e) = socket.try_send_to(&udp_packet, SocketAddr::V6(dst_addr)) + { tracing::error!(?e, "udp send hole punch packet error"); } - tracing::debug!(?dst_addr, "udp forward packet send hole punch packet"); + tracing::debug!( + ?dst_addr, + ?preferred_src, + "udp forward packet send hole punch packet" + ); } else if header.msg_type != UdpPacketType::HolePunch as u8 { let Some(mut conn) = self.sock_map.get_mut(&addr) else { tracing::trace!(?header, "udp forward packet error, connection not found"); return; }; - if let Err(e) = conn.handle_packet_from_remote(zc_packet).await { + if let Err(e) = conn.handle_packet_from_remote(zc_packet) { tracing::trace!(?e, "udp forward packet error"); } } else { @@ -549,7 +650,7 @@ impl UdpTunnelListenerData { let mut buf = BytesMut::new(); loop { match udp_recv_from_socket_forward_task(&socket, &mut buf, true).await { - Ok((zc_packet, addr)) => self.do_forward_one_packet_to_conn(zc_packet, addr).await, + Ok((zc_packet, addr)) => self.do_forward_one_packet_to_conn(zc_packet, addr), Err(e) => { tracing::error!(?e, "udp recv packet error"); break; @@ -705,6 +806,13 @@ impl UdpTunnelConnector { } } + fn should_resend_syn_to_hole_punch_source( + recv_addr: SocketAddr, + expected_addr: SocketAddr, + ) -> bool { + recv_addr == expected_addr + } + async fn wait_sack( socket: &UdpSocket, addr: SocketAddr, @@ -720,12 +828,34 @@ impl UdpTunnelConnector { ) .await??; let zc_packet = get_zcpacket_from_buf(buf.split(), false)?; + let header = zc_packet.udp_tunnel_header().unwrap(); + if header.msg_type == UdpPacketType::HolePunch as u8 { + tracing::debug!(?recv_addr, ?addr, "udp wait sack got hole punch packet"); + if Self::should_resend_syn_to_hole_punch_source(recv_addr, addr) { + let udp_packet = new_syn_packet(conn_id, magic).into_bytes(); + match socket.send_to(&udp_packet, recv_addr).await { + Ok(ret) => { + tracing::debug!(?recv_addr, ?ret, "udp send syn to hole punch source") + } + Err(e) => { + tracing::debug!(?recv_addr, ?e, "udp send syn to hole punch source failed") + } + } + } else { + tracing::debug!( + ?recv_addr, + ?addr, + "ignore hole punch packet from unexpected source" + ); + } + return Err(TunnelError::InvalidPacket( + "got hole punch packet while waiting for sack".to_owned(), + )); + } if recv_addr != addr { tracing::warn!(?recv_addr, ?addr, ?usize, "udp wait sack addr not match"); } - let header = zc_packet.udp_tunnel_header().unwrap(); - if header.conn_id.get() != conn_id { return Err(super::TunnelError::ConnIdNotMatch( header.conn_id.get(), @@ -807,7 +937,7 @@ impl UdpTunnelConnector { match udp_recv_from_socket_forward_task(&socket_clone, &mut buf, false).await { Ok((zc_packet, addr)) => { tracing::trace!(?addr, "connector udp forward task done"); - if let Err(e) = udp_conn.handle_packet_from_remote(zc_packet).await { + if let Err(e) = udp_conn.handle_packet_from_remote(zc_packet) { tracing::trace!(?e, ?addr, "udp forward packet error"); } } @@ -868,6 +998,23 @@ impl UdpTunnelConnector { let udp_packet = new_syn_packet(conn_id, magic).into_bytes(); let ret = socket.send_to(&udp_packet, &addr).await?; tracing::warn!(?udp_packet, ?ret, "udp send syn"); + let resend_task = AbortOnDropHandle::new(tokio::spawn({ + let socket = socket.clone(); + let udp_packet = udp_packet.clone(); + let resend_addr = addr; + async move { + loop { + tokio::time::sleep(Duration::from_millis(200)).await; + match socket.send_to(&udp_packet, &resend_addr).await { + Ok(ret) => tracing::trace!(?ret, ?resend_addr, "udp resend syn"), + Err(e) => { + tracing::debug!(?e, ?resend_addr, "udp resend syn failed"); + break; + } + } + } + } + })); // wait sack let recv_addr = tokio::time::timeout( @@ -875,12 +1022,13 @@ impl UdpTunnelConnector { Self::wait_sack_loop(&socket, addr, conn_id, magic), ) .await??; + drop(resend_task); if recv_addr != addr { tracing::debug!(?recv_addr, ?addr, "udp connect addr not match"); } - self.build_tunnel(socket, addr, conn_id).await + self.build_tunnel(socket, recv_addr, conn_id).await } async fn connect_with_default_bind( @@ -979,9 +1127,43 @@ mod tests { get_interface_name_by_ip, tests::{_tunnel_bench, _tunnel_echo_server, _tunnel_pingpong, wait_for_condition}, }, + packet_def::PacketType, }, }; + fn new_udp_data_packet(conn_id: u32, packet_type: PacketType) -> ZCPacket { + let mut packet = ZCPacket::new_with_payload(b"udp-data").convert_type(ZCPacketType::UDP); + packet.fill_peer_manager_hdr(1, 2, packet_type as u8); + let udp_payload_len = packet.udp_payload().len(); + let header = packet.mut_udp_tunnel_header().unwrap(); + header.conn_id.set(conn_id); + header.msg_type = UdpPacketType::Data as u8; + header.len.set(udp_payload_len as u16); + packet + } + + fn assert_sync_packet_handler(_: fn(&mut UdpConnection, ZCPacket) -> Result<(), TunnelError>) {} + + #[test] + fn hole_punch_source_must_match_connect_addr_before_syn_resend() { + let expected_addr: SocketAddr = "198.51.100.10:11010".parse().unwrap(); + let same_port_different_ip: SocketAddr = "198.51.100.11:11010".parse().unwrap(); + let same_ip_different_port: SocketAddr = "198.51.100.10:11011".parse().unwrap(); + + assert!(UdpTunnelConnector::should_resend_syn_to_hole_punch_source( + expected_addr, + expected_addr + )); + assert!(!UdpTunnelConnector::should_resend_syn_to_hole_punch_source( + same_port_different_ip, + expected_addr + )); + assert!(!UdpTunnelConnector::should_resend_syn_to_hole_punch_source( + same_ip_different_port, + expected_addr + )); + } + #[tokio::test] async fn udp_pingpong() { let listener = UdpTunnelListener::new("udp://0.0.0.0:5556".parse().unwrap()); @@ -989,6 +1171,43 @@ mod tests { _tunnel_pingpong(listener, connector).await; } + #[tokio::test] + async fn udp_connection_handler_uses_sync_nonblocking_ring_delivery() { + assert_sync_packet_handler(UdpConnection::handle_packet_from_remote); + + let socket = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); + let dst_addr = "127.0.0.1:1".parse().unwrap(); + let ring_for_send_udp = Arc::new(RingTunnel::new(8)); + let ring_for_recv_udp = Arc::new(RingTunnel::new(8)); + let (close_event_sender, _close_event_recv) = tokio::sync::mpsc::unbounded_channel(); + let mut conn = UdpConnection::new( + socket, + 7, + dst_addr, + RingSink::new(ring_for_recv_udp), + RingStream::new(ring_for_send_udp), + close_event_sender, + ); + + for _ in 0..16 { + conn.handle_packet_from_remote(new_udp_data_packet(7, PacketType::Data)) + .unwrap(); + } + + let mut got_buffer_full = false; + for _ in 0..16 { + match conn.handle_packet_from_remote(new_udp_data_packet(7, PacketType::Ping)) { + Ok(()) => {} + Err(TunnelError::BufferFull) => { + got_buffer_full = true; + break; + } + Err(e) => panic!("unexpected error: {e:?}"), + } + } + assert!(got_buffer_full); + } + #[tokio::test] async fn udp_bench() { let listener = UdpTunnelListener::new("udp://0.0.0.0:5555".parse().unwrap()); @@ -1212,6 +1431,22 @@ mod tests { .await; } + #[test] + fn v6_hole_punch_packet_preserves_preferred_source_ifindex() { + let dst_addr = "[2001:db8::1]:10001".parse::().unwrap(); + let preferred_src = PreferredIpv6Source { + ip: "2001:db8::2".parse().unwrap(), + ifindex: 42, + }; + + let packet = new_v6_hole_punch_packet(&dst_addr, Some(preferred_src)); + let (parsed_dst_addr, parsed_preferred_src) = + extract_v6_hole_punch_packet(packet.udp_payload()).unwrap(); + + assert_eq!(parsed_dst_addr, dst_addr); + assert_eq!(parsed_preferred_src, Some(preferred_src)); + } + #[tokio::test] async fn test_v6_hole_punch_packet() { let mut lis = UdpTunnelListener::new("udp://[::]:0".parse().unwrap()); @@ -1238,6 +1473,7 @@ mod tests { std::net::SocketAddr::V6(addr_v6) => addr_v6, _ => panic!("Expected an IPv6 address"), }, + None, ) .await .unwrap(); diff --git a/easytier/src/tunnel/udp_src.rs b/easytier/src/tunnel/udp_src.rs new file mode 100644 index 00000000..71f10916 --- /dev/null +++ b/easytier/src/tunnel/udp_src.rs @@ -0,0 +1,210 @@ +use std::{ + io, + net::{Ipv6Addr, SocketAddrV6}, +}; + +use tokio::net::UdpSocket; + +#[cfg(unix)] +pub(crate) fn send_to_with_src_ipv6( + socket: &UdpSocket, + src_ip: Ipv6Addr, + src_ifindex: u32, + dst_addr: SocketAddrV6, + buf: &[u8], +) -> io::Result { + #[cfg(target_env = "ohos")] + { + let _ = (socket, src_ip, src_ifindex, dst_addr, buf); + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "sending UDP with a selected IPv6 source is not supported on OHOS", + )); + } + + #[cfg(not(target_env = "ohos"))] + { + use std::{mem, os::fd::AsRawFd, ptr}; + + use nix::libc; + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + #[cfg(target_os = "android")] + let ipi6_ifindex: libc::c_int = i32::try_from(src_ifindex).map_err(|_| { + io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 source interface index is out of range", + ) + })?; + #[cfg(not(target_os = "android"))] + let ipi6_ifindex: libc::c_uint = src_ifindex; + + let pktinfo = libc::in6_pktinfo { + ipi6_addr: libc::in6_addr { + s6_addr: src_ip.octets(), + }, + ipi6_ifindex, + }; + let mut iov = libc::iovec { + iov_base: buf.as_ptr() as *mut libc::c_void, + iov_len: buf.len(), + }; + let dst_addr = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr)); + let control_len = unsafe { + libc::CMSG_SPACE(mem::size_of::() as libc::c_uint) as usize + }; + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 packet info control buffer is too small", + )); + } + + let mut msg = unsafe { mem::zeroed::() }; + msg.msg_name = dst_addr.as_ptr() as *mut libc::c_void; + msg.msg_namelen = dst_addr.len() as _; + msg.msg_iov = &mut iov; + msg.msg_iovlen = 1; + msg.msg_control = control.0.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control_len as _; + msg.msg_flags = 0; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 packet info control buffer is invalid", + )); + } + (*cmsg).cmsg_level = libc::IPPROTO_IPV6; + (*cmsg).cmsg_type = libc::IPV6_PKTINFO; + (*cmsg).cmsg_len = + libc::CMSG_LEN(mem::size_of::() as libc::c_uint) as _; + ptr::write(libc::CMSG_DATA(cmsg) as *mut libc::in6_pktinfo, pktinfo); + + let ret = libc::sendmsg(socket.as_raw_fd(), &msg, 0); + if ret < 0 { + Err(io::Error::last_os_error()) + } else { + Ok(ret as usize) + } + } + } +} + +#[cfg(windows)] +pub(crate) fn send_to_with_src_ipv6( + socket: &UdpSocket, + src_ip: Ipv6Addr, + src_ifindex: u32, + dst_addr: SocketAddrV6, + buf: &[u8], +) -> io::Result { + use std::{mem, os::windows::io::AsRawSocket, ptr}; + + use windows::{ + Win32::Networking::WinSock::{ + CMSGHDR, IN6_ADDR, IN6_ADDR_0, IN6_PKTINFO, IPPROTO_IPV6, IPV6_PKTINFO, SOCKET, + SOCKET_ERROR, WSABUF, WSAGetLastError, WSAMSG, WSASendMsg, + }, + core::PSTR, + }; + + fn cmsghdr_align(length: usize) -> usize { + (length + mem::align_of::() - 1) & !(mem::align_of::() - 1) + } + + fn cmsgdata_align(length: usize) -> usize { + (length + mem::align_of::() - 1) & !(mem::align_of::() - 1) + } + + fn cmsg_len(length: usize) -> usize { + cmsgdata_align(mem::size_of::()) + length + } + + fn cmsg_space(length: usize) -> usize { + cmsgdata_align(mem::size_of::() + cmsghdr_align(length)) + } + + fn cmsg_data(cmsg: *mut CMSGHDR) -> *mut u8 { + (cmsg as usize + cmsgdata_align(mem::size_of::())) as *mut u8 + } + + #[repr(align(8))] + struct ControlBuffer([u8; 128]); + + let dst = socket2::SockAddr::from(std::net::SocketAddr::V6(dst_addr)); + let mut data = WSABUF { + len: buf.len() as u32, + buf: PSTR(buf.as_ptr() as *mut u8), + }; + let control_len = cmsg_space(mem::size_of::()); + let mut control = ControlBuffer([0u8; 128]); + if control_len > control.0.len() { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "IPv6 packet info control buffer is too small", + )); + } + let mut msg = WSAMSG { + name: dst.as_ptr() as *mut _, + namelen: dst.len(), + lpBuffers: &mut data, + dwBufferCount: 1, + Control: WSABUF { + len: control_len as u32, + buf: PSTR(control.0.as_mut_ptr()), + }, + dwFlags: 0, + }; + + let pktinfo = IN6_PKTINFO { + ipi6_addr: IN6_ADDR { + u: IN6_ADDR_0 { + Byte: src_ip.octets(), + }, + }, + ipi6_ifindex: src_ifindex, + }; + + unsafe { + let cmsg = control.0.as_mut_ptr() as *mut CMSGHDR; + (*cmsg).cmsg_level = IPPROTO_IPV6.0; + (*cmsg).cmsg_type = IPV6_PKTINFO; + (*cmsg).cmsg_len = cmsg_len(mem::size_of::()); + ptr::write(cmsg_data(cmsg) as *mut IN6_PKTINFO, pktinfo); + msg.Control.len = control_len as u32; + + let mut sent = 0; + let ret = WSASendMsg( + SOCKET(socket.as_raw_socket() as usize), + &msg, + 0, + Some(&mut sent), + None, + None, + ); + if ret == SOCKET_ERROR { + return Err(io::Error::from_raw_os_error(WSAGetLastError().0)); + } + Ok(sent as usize) + } +} + +#[cfg(not(any(unix, windows)))] +pub(crate) fn send_to_with_src_ipv6( + _socket: &UdpSocket, + _src_ip: Ipv6Addr, + _src_ifindex: u32, + _dst_addr: SocketAddrV6, + _buf: &[u8], +) -> io::Result { + Err(io::Error::new( + io::ErrorKind::Unsupported, + "sending UDP with a selected IPv6 source is not supported on this platform", + )) +}