From 2dc921b2d521e788c287a4e5a2869f80cabed5f1 Mon Sep 17 00:00:00 2001 From: DuProcess <273172371+DuProcess@users.noreply.github.com> Date: Thu, 16 Jul 2026 20:46:21 -0400 Subject: [PATCH] Bind client UDP sockets by peer family --- src/bin/dosh-client.rs | 10 +++++----- src/client.rs | 5 ++--- src/udp.rs | 28 ++++++++++++++++++++++++++-- 3 files changed, 33 insertions(+), 10 deletions(-) diff --git a/src/bin/dosh-client.rs b/src/bin/dosh-client.rs index a823338..72eb0a6 100644 --- a/src/bin/dosh-client.rs +++ b/src/bin/dosh-client.rs @@ -37,8 +37,8 @@ use dosh::transport::{ TransportEvent, }; use dosh::udp::{ - is_transient_udp_error, is_transient_udp_send_error, recv_udp_retrying_transient, - send_udp_retrying_transient, + bind_udp_for_peer, is_transient_udp_error, is_transient_udp_send_error, + recv_udp_retrying_transient, send_udp_retrying_transient, }; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; @@ -451,7 +451,7 @@ async fn main() -> Result<()> { }; let cache_requested_env = requested_env(&config, &host, resolved_ssh_config.as_ref()); - let socket = UdpSocket::bind("0.0.0.0:0").await?; + let socket = bind_udp_for_peer(resolve_addr(&target_udp_host, dosh_port)?).await?; if let Some(mut cached) = credential.clone() { cached.udp_host = target_udp_host.clone(); @@ -1955,7 +1955,7 @@ async fn run_doctor_for_host( "[info] forwarding requested by config: agent={} default_session={}", config.forward_agent, config.default_session ); - let socket = UdpSocket::bind("0.0.0.0:0").await?; + let socket = bind_udp_for_peer(resolve_addr(&udp_host, dosh_port)?).await?; match try_native_auth_check( &socket, config, @@ -2393,7 +2393,7 @@ async fn run_proxy_stdio_command(config: &dosh::config::ClientConfig, args: &Arg let target_udp_host = selected_udp_host(requested_udp_host.as_deref(), &server, &resolved_ssh_config)?; let dosh_port = args.dosh_port.or(host.port).unwrap_or(config.dosh_port); - let socket = UdpSocket::bind("0.0.0.0:0").await?; + let socket = bind_udp_for_peer(resolve_addr(&target_udp_host, dosh_port)?).await?; let session = args .session .clone() diff --git a/src/client.rs b/src/client.rs index 21808cd..bc23084 100644 --- a/src/client.rs +++ b/src/client.rs @@ -13,12 +13,11 @@ use crate::protocol::{ }; use crate::ssh_agent; use crate::transport::{DoshTransport, SessionRole, SessionTransportConfig, TransportConfig}; -use crate::udp::{recv_udp_retrying_transient, send_udp_retrying_transient}; +use crate::udp::{bind_udp_for_peer, recv_udp_retrying_transient, send_udp_retrying_transient}; use anyhow::{Context, Result, anyhow, bail}; use std::net::{SocketAddr, ToSocketAddrs}; use std::path::PathBuf; use std::time::{Duration, SystemTime, UNIX_EPOCH}; -use tokio::net::UdpSocket; #[derive(Debug, Clone)] pub struct DoshClient { @@ -189,7 +188,7 @@ impl DoshClientBuilder { }) }) .collect::>>()?; - let socket = UdpSocket::bind("0.0.0.0:0").await?; + let socket = bind_udp_for_peer(peer_addr).await?; let (client_secret, client_public) = generate_native_ephemeral(); let hello = NativeClientHello { protocol_version: native::NATIVE_PROTOCOL_VERSION, diff --git a/src/udp.rs b/src/udp.rs index 082df41..ef19585 100644 --- a/src/udp.rs +++ b/src/udp.rs @@ -1,6 +1,6 @@ //! UDP error classification shared by Dosh clients, servers, and embedders. -use std::net::SocketAddr; +use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr}; use std::time::{Duration, Instant}; use tokio::net::UdpSocket; @@ -17,6 +17,17 @@ pub fn is_transient_udp_send_error(err: &std::io::Error) -> bool { is_transient_udp_error(err) } +pub fn unspecified_bind_addr_for_peer(peer: SocketAddr) -> SocketAddr { + match peer { + SocketAddr::V4(_) => SocketAddr::from((Ipv4Addr::UNSPECIFIED, 0)), + SocketAddr::V6(_) => SocketAddr::from((Ipv6Addr::UNSPECIFIED, 0)), + } +} + +pub async fn bind_udp_for_peer(peer: SocketAddr) -> std::io::Result { + UdpSocket::bind(unspecified_bind_addr_for_peer(peer)).await +} + pub async fn send_udp_retrying_transient( socket: &UdpSocket, packet: &[u8], @@ -106,8 +117,9 @@ fn is_transient_udp_os_error(_code: i32) -> bool { mod tests { use super::{ is_transient_udp_error, is_transient_udp_send_error, recv_udp_retrying_transient, - send_udp_retrying_transient, + send_udp_retrying_transient, unspecified_bind_addr_for_peer, }; + use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr}; use std::time::Duration; use tokio::net::UdpSocket; @@ -127,6 +139,18 @@ mod tests { assert!(!is_transient_udp_send_error(&fatal)); } + #[test] + fn peer_bind_addr_matches_peer_address_family() { + assert_eq!( + unspecified_bind_addr_for_peer(SocketAddr::from((Ipv4Addr::LOCALHOST, 50000))), + SocketAddr::from((Ipv4Addr::UNSPECIFIED, 0)) + ); + assert_eq!( + unspecified_bind_addr_for_peer(SocketAddr::from((Ipv6Addr::LOCALHOST, 50000))), + SocketAddr::from((Ipv6Addr::UNSPECIFIED, 0)) + ); + } + #[tokio::test] async fn retrying_udp_send_and_receive_round_trip() { let receiver = UdpSocket::bind("127.0.0.1:0").await.unwrap();