diff --git a/book/src/SUMMARY.md b/book/src/SUMMARY.md index ca9fcbf..0a33292 100644 --- a/book/src/SUMMARY.md +++ b/book/src/SUMMARY.md @@ -16,7 +16,7 @@ - [Exposing your first service](./exposing_your_first_service.md) - [Local forwarding](./local_forwarding.md) - [Custom domains](./custom_domains.md) -- [UDP-over-TCP](./udp_over_tcp.md) +- [UDP-over-TCP (experimental)](./udp_over_tcp.md) - [Advanced options](./advanced_options.md) # Reference diff --git a/book/src/udp_over_tcp.md b/book/src/udp_over_tcp.md index 48bc6c0..9b4d6ef 100644 --- a/book/src/udp_over_tcp.md +++ b/book/src/udp_over_tcp.md @@ -1,4 +1,4 @@ -# UDP-over-TCP +# UDP-over-TCP (experimental) Sandhole has experimental support for UDP over SSH, with a thin TCP-based protocol. diff --git a/src/ssh/auth.rs b/src/ssh/auth.rs index 2bd6c89..9d20dec 100644 --- a/src/ssh/auth.rs +++ b/src/ssh/auth.rs @@ -16,8 +16,8 @@ use tokio_util::sync::CancellationToken; use crate::{ admin::interface::AdminInterface, connection_handler::ConnectionHttpData, - droppable_handle::DroppableHandle, ip::IpFilter, quota::TokenHolder, ssh::FingerprintFn, - sock_addr_alias::SockAddrAlias, + droppable_handle::DroppableHandle, ip::IpFilter, quota::TokenHolder, + sock_addr_alias::SockAddrAlias, ssh::FingerprintFn, }; pub(crate) struct ProxyAutoCancellation { diff --git a/src/ssh/exec.rs b/src/ssh/exec.rs index 2e9b0ed..36afa95 100644 --- a/src/ssh/exec.rs +++ b/src/ssh/exec.rs @@ -8,8 +8,8 @@ use rustls_pki_types::DnsName; use crate::{ SandholeServer, admin::interface::AdminInterface, - ssh::{AuthenticatedData, ServerHandlerSender, auth::UserSessionRestriction}, sock_addr_alias::SockAddrAlias, + ssh::{AuthenticatedData, ServerHandlerSender, auth::UserSessionRestriction}, }; #[bitflags] diff --git a/src/udp.rs b/src/udp.rs index 8723300..c92a04e 100644 --- a/src/udp.rs +++ b/src/udp.rs @@ -1,4 +1,9 @@ -use std::{collections::HashSet, mem::size_of, net::IpAddr, sync::Arc}; +use std::{ + collections::HashSet, + mem::size_of, + net::{IpAddr, SocketAddr}, + sync::Arc, +}; use crate::{ connection_handler::ConnectionHandler, @@ -6,7 +11,7 @@ use crate::{ droppable_handle::DroppableHandle, ip::IpFilter, reactor::UdpReactor, - ssh::connection_handler::SshTunnelHandler, + ssh::connection_handler::{SshChannel, SshTunnelHandler}, telemetry::{TELEMETRY_COUNTER_UDP_CONNECTIONS, TELEMETRY_KEY_PORT}, udp_listener::get_udp_socket, }; @@ -14,11 +19,10 @@ use ahash::RandomState; use bon::Builder; use color_eyre::eyre::Context; use dashmap::DashMap; -use futures_util::pin_mut; use metrics::counter; use tokio::{ - io::{AsyncReadExt, AsyncWriteExt}, - select, + io::{AsyncReadExt, AsyncWriteExt, WriteHalf}, + sync::Mutex, }; pub const MAX_PACKET_SIZE: usize = size_of::() + u16::MAX as usize; @@ -33,6 +37,9 @@ fn datagram_buffer() -> Box<[u8; MAX_PACKET_SIZE]> { Box::new([0u8; MAX_PACKET_SIZE]) } +// Type for a UDP socket with the write and read halves. +type UdpSocketHandler = (Arc>>, DroppableHandle<()>); + // Service that handles creating UDP sockets for reverse forwarding connections. #[derive(Builder)] pub(crate) struct UdpHandler { @@ -40,7 +47,10 @@ pub(crate) struct UdpHandler { listen_address: IpAddr, // Map containing spawned tasks of connections for each socket. #[builder(skip = DashMap::default())] - sockets: DashMap, RandomState>, + tasks: DashMap, RandomState>, + // Map linking a stateless socket to its underlying SSH channel. + #[builder(skip = DashMap::default())] + sockets: DashMap<(u16, SocketAddr), UdpSocketHandler, RandomState>, // Connection map to assign a tunneling service for each incoming connection. conn_manager: Arc, UdpReactor>>, // Service that identifies whether to allow or block a given IP address. @@ -58,20 +68,18 @@ pub(crate) trait UdpPortHandler { impl UdpPortHandler for Arc { // Create a UDP listener on the given port. async fn create_port_socket(&self, port: u16) -> color_eyre::Result { - if self.sockets.contains_key(&port) { + if self.tasks.contains_key(&port) { return Ok(port); } // Check if we're able to bind to the given address and port. - let socket = get_udp_socket((self.listen_address, port))?; + let socket = Arc::new(get_udp_socket((self.listen_address, port))?); let port = socket .local_addr() .with_context(|| "Missing local address when binding port")? .port(); let clone = Arc::clone(self); - let listen_address = self.listen_address; // Start task that will listen to incoming connections. let join_handle = DroppableHandle(tokio::spawn(async move { - let mut socket = socket; loop { let mut buf = datagram_buffer(); match socket.recv_from(buf.as_mut()).await { @@ -82,6 +90,30 @@ impl UdpPortHandler for Arc { tracing::info!(%address, "Rejecting UDP connection: IP not allowed."); continue; } + + // Prepare data to be sent to SSH channel + let mut read_buf = datagram_buffer(); + read_buf[..size_of::()] + .copy_from_slice(&(len as u16).to_be_bytes()[..]); + read_buf[size_of::()..size_of::() + len] + .copy_from_slice(&buf[..len]); + + // Check for an existing SSH channel + if let Some(entry) = clone.sockets.get(&(port, address)) { + let channel = Arc::clone(&entry.value().0); + if channel + .lock() + .await + .write_all(&read_buf[..size_of::() + len]) + .await + .is_ok() + { + continue; + } else { + clone.sockets.remove(&(port, address)); + } + } + // Get the handler for this port if let Some(handler) = clone.conn_manager.get(&port, ip) && let Ok(channel) = handler.tunneling_channel(ip, address.port()).await @@ -100,117 +132,67 @@ impl UdpPortHandler for Arc { .into_bytes(), ); } - let socket = std::mem::replace( - &mut socket, - get_udp_socket((listen_address, port)) - .expect("should re-create UDP socket"), - ); - if let Err(error) = socket.connect(address).await { + + let udp_write = Arc::clone(&socket); + let (mut ssh_read, mut ssh_write) = tokio::io::split(channel); + + if let Err(error) = ssh_write + .write_all(&read_buf[..size_of::() + len]) + .await + { #[cfg(not(coverage_nightly))] - tracing::error!(%port, %error, "Error connecting UDP socket."); + tracing::warn!(%port, %error, "Error sending UDP datagram to SSH channel."); continue; } - let mut read_buf = datagram_buffer(); - - // Prepare already consumed data to be sent to SSH channel - *read_buf[..size_of::()] - .as_mut_array() - .expect("length checked") = (len as u16).to_be_bytes(); - read_buf[size_of::()..size_of::() + len] - .copy_from_slice(&buf[..len]); - - tokio::spawn(async move { - let udp_read = Arc::new(socket); - let udp_write = udp_read.clone(); - let (mut ssh_read, mut ssh_write) = tokio::io::split(channel); - - let udp2ssh = async move { - let mut buf = read_buf; - if let Err(error) = - ssh_write.write_all(&buf[..size_of::() + len]).await - { - #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error writing to SSH channel for UDP."); - return; - } - loop { - match udp_read - .recv(&mut buf.as_mut()[size_of::()..]) - .await - { - Ok(len) => { - *buf[..size_of::()] - .as_mut_array() - .expect("length checked") = - (len as u16).to_be_bytes(); - if let Err(error) = ssh_write - .write_all(&buf[..size_of::() + len]) - .await - { - #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error writing to SSH channel for UDP."); - break; - }; - } - Err(error) => { + + let read_handle = DroppableHandle(tokio::spawn(async move { + let mut write_buf = datagram_buffer(); + loop { + match ssh_read.read_u16().await { + Ok(len) => { + if let Err(error) = ssh_read + .read_exact(&mut write_buf[..len as usize]) + .await + { #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error reading from UDP socket."); + tracing::warn!(%port, %error, "Error reading UDP datagram from SSH channel."); break; - } - } - } - }; - let ssh2udp = async move { - let mut buf = datagram_buffer(); - loop { - match ssh_read.read_u16().await { - Ok(len) => { - if let Err(error) = ssh_read - .read_exact(&mut buf[..len as usize]) - .await - { - #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error reading UDP datagram from SSH channel."); - break; - } else if let Err(error) = - udp_write.send(&mut buf[..len as usize]).await - { - #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error reading from SSH channel for UDP."); - break; - } - } - Err(error) => { + } else if let Err(error) = udp_write + .send_to(&write_buf[..len as usize], address) + .await + { #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error reading UDP datagram size from SSH channel."); + tracing::warn!(%port, %error, "Error reading from SSH channel for UDP."); break; } } + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error reading UDP datagram size from SSH channel."); + break; + } } - }; - - pin_mut!(udp2ssh); - pin_mut!(ssh2udp); - - select! { - _ = udp2ssh => {} - _ = ssh2udp => {} } - }); + })); + + clone.sockets.insert( + (port, address), + (Arc::new(Mutex::new(ssh_write)), read_handle), + ); } } Err(error) => { #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error getting remote connection on UDP port.") + tracing::warn!(%port, %error, "Error getting remote connection on UDP port."); } } } })); - self.sockets.insert(port, join_handle); + self.tasks.insert(port, join_handle); Ok(port) } - // Create a TCP listener on a random open port, returning the port number. + // Create a UDP listener on a random open port, returning the port number. async fn get_free_port(&self) -> color_eyre::Result { // By passing 0 to create_port_listener, the OS will choose a port for us. self.create_port_socket(0).await @@ -221,9 +203,9 @@ impl UdpPortHandler for Arc { // Find the ports listening to the localhost address let mut ports: HashSet = ports.into_iter().collect(); // Remove any socket tasks not in the list of localhost port - self.sockets.retain(|port, _| ports.contains(port)); + self.tasks.retain(|port, _| ports.contains(port)); // Find the list of new ports - ports.retain(|port| !self.sockets.contains_key(port)); + ports.retain(|port| !self.tasks.contains_key(port)); if !ports.is_empty() { let clone = Arc::clone(self); // Create port listeners for the new ports @@ -231,7 +213,7 @@ impl UdpPortHandler for Arc { for port in ports.into_iter() { if let Err(error) = clone.create_port_socket(port).await { #[cfg(not(coverage_nightly))] - tracing::error!(%port, %error, "Failed to create listener for TCP port."); + tracing::error!(%port, %error, "Failed to create listener for UDP port."); } } }); diff --git a/tests/integration/udp_ip_connections_limit.rs b/tests/integration/udp_ip_connections_limit.rs index f167e23..402a994 100644 --- a/tests/integration/udp_ip_connections_limit.rs +++ b/tests/integration/udp_ip_connections_limit.rs @@ -111,7 +111,7 @@ async fn udp_ip_connections_limit() { .await .expect("UDP connection failed"); udp_socket - .connect(format!("127.0.0.1:12345")) + .connect("127.0.0.1:12345".to_string()) .await .unwrap(); udp_socket.send(b"0123456789").await.unwrap(); @@ -134,7 +134,7 @@ async fn udp_ip_connections_limit() { .await .expect("UDP connection failed"); udp_socket - .connect(format!("127.0.0.1:12345")) + .connect("127.0.0.1:12345".to_string()) .await .unwrap(); udp_socket.send(b"0123456789").await.unwrap(); @@ -153,7 +153,7 @@ async fn udp_ip_connections_limit() { let udp_socket = UdpSocket::bind("[::1]:0") .await .expect("UDP connection failed"); - udp_socket.connect(format!("[::1]:12345")).await.unwrap(); + udp_socket.connect("[::1]:12345".to_string()).await.unwrap(); udp_socket.send(b"0123456789").await.unwrap(); assert!(started.elapsed() < Duration::from_secs(5)); let mut data = [0u8; 16]; diff --git a/tests/integration/udp_multi_stream.rs b/tests/integration/udp_multi_stream.rs index eae5f91..85c8649 100644 --- a/tests/integration/udp_multi_stream.rs +++ b/tests/integration/udp_multi_stream.rs @@ -133,7 +133,7 @@ async fn udp_multi_stream() { .expect("UDP connection failed"), ); udp_socket_read - .connect(format!("127.0.0.1:12345")) + .connect("127.0.0.1:12345".to_string()) .await .unwrap(); let udp_socket_write = Arc::clone(&udp_socket_read); diff --git a/tests/integration/udp_rate_limit.rs b/tests/integration/udp_rate_limit.rs index afc9ae1..657d5e5 100644 --- a/tests/integration/udp_rate_limit.rs +++ b/tests/integration/udp_rate_limit.rs @@ -119,13 +119,19 @@ async fn udp_rate_limit() { .await .expect("UDP connection failed"); udp_socket - .connect(format!("127.0.0.1:12345")) + .connect("127.0.0.1:12345".to_string()) .await .unwrap(); let start = Instant::now(); udp_socket.send(&data[..]).await.unwrap(); let mut buf = [0u8; 32]; - assert_eq!(udp_socket.recv(&mut buf).await.unwrap(), 2); + assert_eq!( + timeout(Duration::from_secs(5), udp_socket.recv(&mut buf)) + .await + .unwrap() + .unwrap(), + 2 + ); let elapsed = start.elapsed(); assert_eq!(&buf[..2], b"OK"); assert!( diff --git a/udp_over_tcp.py b/udp_over_tcp.py index acdd1bb..335c87a 100644 --- a/udp_over_tcp.py +++ b/udp_over_tcp.py @@ -58,9 +58,8 @@ class UdpProxyProtocol(asyncio.Protocol): # -- Remote forwarding -- class TcpServerProtocol(TcpProxyProtocol): - def __init__(self, udp_address, udp_port): + def __init__(self, udp_address): self.udp_address = udp_address - self.udp_port = udp_port self.task = None self.tcp_transport = None self.on_connection_lost = asyncio.get_running_loop().create_future() @@ -80,7 +79,7 @@ class TcpServerProtocol(TcpProxyProtocol): loop = asyncio.get_running_loop() udp_transport, _ = await loop.create_datagram_endpoint( lambda: UdpClientProtocol(self.tcp_transport), - remote_addr=(self.udp_address, self.udp_port), + remote_addr=self.udp_address, ) self.udp_transport = udp_transport @@ -99,6 +98,9 @@ class UdpClientProtocol(UdpProxyProtocol): self.tcp_transport = tcp_transport self.buffered_data = [] + def error_received(self, exc): + pass + # -- Local forwarding -- class TcpClientProtocol(TcpProxyProtocol): @@ -183,7 +185,7 @@ async def main(): udp_transport.close() else: server = await loop.create_server( - lambda: TcpServerProtocol(args.udp_address or "127.0.0.1", args.udp_port), + lambda: TcpServerProtocol((args.udp_address or "127.0.0.1", args.udp_port)), args.tcp_address or "0.0.0.0", args.tcp_port, )