diff --git a/.github/workflows/nix.yml b/.github/workflows/nix.yml index afc4f57..ad59d77 100644 --- a/.github/workflows/nix.yml +++ b/.github/workflows/nix.yml @@ -45,4 +45,5 @@ jobs: CACHIX_NAME: ${{ vars.CACHIX_NAME }} CACHIX_AUTH_TOKEN: ${{ secrets.CACHIX_AUTH_TOKEN }} run: | - nix build --no-link --print-out-paths | cachix push ${CACHIX_NAME} + nix build .#sandhole --no-link --print-out-paths | cachix push ${CACHIX_NAME} + nix build .#udp_over_tcp --no-link --print-out-paths | cachix push ${CACHIX_NAME} diff --git a/Cargo.lock b/Cargo.lock index 1182379..48511b8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4503,9 +4503,9 @@ dependencies = [ "rustls-pki-types", "rustls-webpki", "rustrict", + "sandhole_socket", "serde", "serde_json", - "socket2 0.6.3", "sysinfo", "test-log", "thiserror 2.0.18", @@ -4524,6 +4524,27 @@ dependencies = [ "webpki-roots 1.0.7", ] +[[package]] +name = "sandhole_socket" +version = "0.1.0" +dependencies = [ + "color-eyre", + "socket2 0.6.3", + "tokio", +] + +[[package]] +name = "sandhole_udp_over_tcp" +version = "0.1.0" +dependencies = [ + "ahash", + "clap", + "color-eyre", + "dashmap", + "sandhole_socket", + "tokio", +] + [[package]] name = "schannel" version = "0.1.28" diff --git a/Cargo.toml b/Cargo.toml index e251a36..5fee9d5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,6 +13,18 @@ keywords = ["ssh", "proxy", "reverse-proxy", "tunnel", "hole-punching"] categories = ["network-programming", "web-programming", "authentication"] exclude = [".github", "book", "docker-compose-example", "tests/data"] +[workspace] +members = ["sandhole_socket", "udp_over_tcp"] + +[workspace.dependencies] +ahash = "0.8.12" +clap = { version = "4.6.1", features = ["derive", "string"] } +color-eyre = "0.6.5" +dashmap = "6.1.0" +sandhole_socket = { version = "0.1.0", path = "sandhole_socket" } +socket2 = "0.6.3" +tokio = { version = "1.52.3", features = ["full"] } + [features] default = ["acme", "duper", "login", "prometheus", "rustrict"] acme = ["dep:rustls-acme"] @@ -21,7 +33,7 @@ login = ["dep:reqwest"] prometheus = ["dep:metrics-exporter-prometheus"] [dependencies] -ahash = "0.8.12" +ahash = { workspace = true } async-speed-limit = { version = "0.4.2", features = ["tokio"] } aws-lc-rs = "1.17.0" axum = { version = "0.8.9", default-features = false } @@ -30,11 +42,11 @@ bon = "3.9.1" bytes = "1.11.1" bytesize = "2.3.1" chrono = "0.4.44" -clap = { version = "4.6.1", features = ["derive", "string"] } +clap = { workspace = true } clap_complete = "4.6.5" -color-eyre = "0.6.5" +color-eyre = { workspace = true } crossterm = { version = "0.29.0", default-features = false } -dashmap = "6.1.0" +dashmap = { workspace = true } deadpool = { version = "0.13.0", default-features = false, features = ["rt_tokio_1", "unmanaged"] } enumflags2 = "0.7.12" futures-util = { version = "0.3.32", default-features = false, features = [ @@ -83,14 +95,14 @@ rustls-acme = { version = "0.15.2", optional = true, default-features = false, f rustls-pki-types = "1.14.1" rustls-webpki = "0.103.13" rustrict = { version = "0.7.38", optional = true, features = ["customize"] } +sandhole_socket = { workspace = true } serde = { version = "1.0.228", default-features = false, features = [ "serde_derive", ] } serde_json = "1.0.149" -socket2 = "0.6.3" sysinfo = "0.38.4" # 0.39.x requires MSRV bump to 1.95.0 thiserror = "2.0.18" -tokio = { version = "1.52.3", features = ["full"] } +tokio = { workspace = true } tokio-rustls = "0.26.4" tokio-stream = "0.1.18" tokio-util = { version = "0.7.18", features = ["net"] } diff --git a/book/src/nixos.md b/book/src/nixos.md index 4882de1..48b0fc4 100644 --- a/book/src/nixos.md +++ b/book/src/nixos.md @@ -118,8 +118,8 @@ in { # Change these from `sandhole.com.br` to your domain domains = [ "sandhole.com.br" "*.sandhole.com.br" ]; - fullchain_output_file = "${certificatesDirectory}/sandhole.com.br/fullchain.pem"; - key_output_file = "${certificatesDirectory}/sandhole.com.br/privkey.pem"; + fullchain_output_file = "${certificates-directory}/sandhole.com.br/fullchain.pem"; + key_output_file = "${certificates-directory}/sandhole.com.br/privkey.pem"; } ]; } @@ -202,17 +202,15 @@ You can then connect services with the provided keys. For example, to use a Vaul ## Binary caching -In order to avoid re-building Sandhole for each update, you can use either of the Sandhole binary caches. In `configuration.nix`: +In order to avoid re-building Sandhole for each update, you can use Sandhole's binary cache. In `configuration.nix`: ```nix nix.settings = { substituters = [ "https://sandhole.cachix.org" - "https://cache.garnix.io" ]; trusted-public-keys = [ "sandhole.cachix.org-1:cZadr6kgjQcRvsr++Nv9kgtMOrbLahiZBpuI9WpIXvA=" - "cache.garnix.io:CTFPyKSLcx5RMJKfLo5EEPUObbA78b0YQ2DTCJXqr9g=" ]; }; ``` diff --git a/nix/checks.nix b/nix/checks.nix index 57ca82b..e7047d0 100644 --- a/nix/checks.nix +++ b/nix/checks.nix @@ -4,12 +4,13 @@ craneLib, pkgs, sandhole, - sandhole-no-default-features, + sandhole-no_default_features, src, + udp_over_tcp, ... }: { - inherit sandhole sandhole-no-default-features; + inherit sandhole sandhole-no_default_features udp_over_tcp; sandhole-clippy = craneLib.cargoClippy ( commonArgs diff --git a/nix/default.nix b/nix/default.nix index ce9a8e9..923dbff 100644 --- a/nix/default.nix +++ b/nix/default.nix @@ -59,18 +59,40 @@ let } ); - sandhole-no-default-features = sandhole.overrideAttrs { + sandhole-no_default_features = sandhole.overrideAttrs { cargoExtraArgs = "--locked --no-default-features"; }; + + udp_over_tcp = craneLib.buildPackage ( + commonArgs + // { + inherit cargoArtifacts; + inherit (craneLib.crateNameFromCargoToml { cargoToml = ../udp_over_tcp/Cargo.toml; }) + pname + version + ; + doCheck = false; + cargoExtraArgs = "-p sandhole_udp_over_tcp"; + meta = { + name = "sandhole_udp_over_tcp"; + description = "Proxy UDP traffic for Sandhole via SSH"; + homepage = "https://sandhole.com.br"; + license = lib.licenses.mit; + mainProgram = "sandhole_udp_over_tcp"; + platforms = lib.platforms.linux ++ lib.platforms.darwin; + }; + } + ); in { - inherit sandhole sandhole-no-default-features; + inherit sandhole sandhole-no_default_features; packages = import ./packages.nix { inherit pkgs sandhole - sandhole-no-default-features + sandhole-no_default_features + udp_over_tcp ; }; @@ -81,7 +103,8 @@ in craneLib pkgs sandhole - sandhole-no-default-features + sandhole-no_default_features + udp_over_tcp src ; }; diff --git a/nix/packages.nix b/nix/packages.nix index 0d44937..a579bb7 100644 --- a/nix/packages.nix +++ b/nix/packages.nix @@ -1,7 +1,8 @@ { pkgs, sandhole, - sandhole-no-default-features, + sandhole-no_default_features, + udp_over_tcp, ... }: let @@ -26,20 +27,9 @@ let }; in { - inherit sandhole sandhole-no-default-features; + inherit sandhole sandhole-no_default_features udp_over_tcp; default = sandhole; - udp_over_tcp = pkgs.python313Packages.buildPythonApplication (finalAttrs: { - name = "udp_over_tcp"; - pyproject = false; - doCheck = false; - dontUnpack = true; - installPhase = '' - install -Dm755 "${../udp_over_tcp.py}" "$out/bin/${finalAttrs.name}" - ''; - meta.mainProgram = finalAttrs.name; - }); - _docs = (pkgs.nixosOptionsDoc { options = removeAttrs evalOptions.options [ "_module" ]; diff --git a/sandhole_socket/Cargo.toml b/sandhole_socket/Cargo.toml new file mode 100644 index 0000000..2d16874 --- /dev/null +++ b/sandhole_socket/Cargo.toml @@ -0,0 +1,18 @@ +[package] +name = "sandhole_socket" +version = "0.1.0" +edition = "2024" +rust-version = "1.88.0" +description = "Sandhole utilities for handling TCP and UDP sockets." +repository = "https://github.com/EpicEric/sandhole" +homepage = "https://sandhole.com.br/" +license = "MIT" +authors = ["Eric Rodrigues Pires "] +readme = "../README.md" +keywords = ["ssh", "proxy", "reverse-proxy", "tcp", "udp"] +categories = ["network-programming"] + +[dependencies] +color-eyre = { workspace = true } +socket2 = { workspace = true } +tokio = { workspace = true } diff --git a/sandhole_socket/src/lib.rs b/sandhole_socket/src/lib.rs new file mode 100644 index 0000000..fd5d6cb --- /dev/null +++ b/sandhole_socket/src/lib.rs @@ -0,0 +1,3 @@ +pub mod tcp_listener; +pub mod udp_listener; +pub mod udp_over_tcp; diff --git a/src/tcp_listener.rs b/sandhole_socket/src/tcp_listener.rs similarity index 86% rename from src/tcp_listener.rs rename to sandhole_socket/src/tcp_listener.rs index 00e1c7d..996cc2e 100644 --- a/src/tcp_listener.rs +++ b/sandhole_socket/src/tcp_listener.rs @@ -3,9 +3,9 @@ use std::{io, net::ToSocketAddrs}; use socket2::{Domain, Socket, Type}; use tokio::net::TcpListener; -// Create an async TCP listener with Nagle's algorithm disabled -// and any necessary configurations for dualstack. -pub(crate) fn get_tcp_listener(addr: A) -> io::Result { +/// Create an async TCP listener with Nagle's algorithm disabled +/// and any necessary configurations for dualstack. +pub fn get_tcp_listener(addr: A) -> io::Result { let addr = addr.to_socket_addrs()?.next().ok_or_else(|| { io::Error::new( io::ErrorKind::InvalidInput, diff --git a/src/udp_listener.rs b/sandhole_socket/src/udp_listener.rs similarity index 50% rename from src/udp_listener.rs rename to sandhole_socket/src/udp_listener.rs index 62e5968..5f405ba 100644 --- a/src/udp_listener.rs +++ b/sandhole_socket/src/udp_listener.rs @@ -3,8 +3,8 @@ use std::{io, net::ToSocketAddrs}; use socket2::{Domain, Socket, Type}; use tokio::net::UdpSocket; -// Create an async UDP listener. -pub(crate) fn get_udp_socket(addr: A) -> io::Result { +/// Create an async UDP listener. +pub fn get_udp_socket(addr: A) -> io::Result { let addr = addr.to_socket_addrs()?.next().ok_or_else(|| { io::Error::new( io::ErrorKind::InvalidInput, @@ -24,9 +24,18 @@ pub(crate) fn get_udp_socket(addr: A) -> io::Result socket.set_only_v6(false)?; } - socket.set_reuse_address(true)?; + // On platforms with Berkeley-derived sockets, this allows to quickly + // rebind a socket, without needing to wait for the OS to clean up the + // previous one. + // + // On Windows, this allows rebinding sockets which are actively in use, + // which allows “socket hijacking”, so we explicitly don't set it here. + // https://docs.microsoft.com/en-us/windows/win32/winsock/using-so-reuseaddr-and-so-exclusiveaddruse #[cfg(not(windows))] - socket.set_reuse_port(true)?; + { + socket.set_reuse_address(true)?; + socket.set_reuse_port(true)?; + } socket.bind(&addr.into())?; diff --git a/sandhole_socket/src/udp_over_tcp.rs b/sandhole_socket/src/udp_over_tcp.rs new file mode 100644 index 0000000..ed36241 --- /dev/null +++ b/sandhole_socket/src/udp_over_tcp.rs @@ -0,0 +1,43 @@ +use std::{mem::size_of, pin::Pin}; + +use color_eyre::eyre::Context; +use tokio::io::{AsyncRead, AsyncReadExt}; + +/// Maximum buffer size required to include the UDP datagram + a length header. +pub const MAX_PACKET_SIZE: usize = size_of::() + u16::MAX as usize; + +/// Creates and returns a buffer on the heap with enough space to contain any possible +/// UDP datagram. +/// +/// This is put on the heap and in a separate function to avoid the 64k buffer from ending +/// up on the stack and blowing up the size of the futures using it. +#[inline] +pub fn datagram_buffer() -> Box<[u8; MAX_PACKET_SIZE]> { + Box::new([0u8; MAX_PACKET_SIZE]) +} + +/// Serialize an UDP datagram into the provided buffer. +#[inline] +pub fn serialize_datagram(buf: &mut [u8], datagram: &[u8]) -> usize { + let len = datagram.len(); + buf[..size_of::()].copy_from_slice(&(len as u16).to_be_bytes()[..]); + buf[size_of::()..size_of::() + len].copy_from_slice(&datagram[..len]); + size_of::() + len +} + +/// Deserialize a source of bytes into an UDP datagram at the provided buffer. +#[inline] +pub async fn deserialize_datagram( + buf: &mut [u8], + source: &mut Pin<&mut R>, +) -> color_eyre::Result { + let len = source + .read_u16() + .await + .with_context(|| "Couldn't read UDP datagram size from source")?; + source + .read_exact(&mut buf[..len as usize]) + .await + .with_context(|| "Couldn't read UDP data size from source")?; + Ok(len as usize) +} diff --git a/src/entrypoint.rs b/src/entrypoint.rs index 829052c..13044ef 100644 --- a/src/entrypoint.rs +++ b/src/entrypoint.rs @@ -25,6 +25,7 @@ use rustls::ServerConfig; #[cfg(feature = "acme")] use rustls_acme::acme::ACME_TLS_ALPN_NAME; use rustls_pki_types::ServerName; +use sandhole_socket::tcp_listener::get_tcp_listener; use sysinfo::{CpuRefreshKind, MemoryRefreshKind, Networks, RefreshKind, System}; #[cfg_attr(not(feature = "acme"), allow(unused_imports))] use tokio::io::AsyncWriteExt; @@ -63,7 +64,6 @@ use crate::{ reactor::{AliasReactor, HttpReactor, SniReactor, SshReactor, TcpReactor, UdpReactor}, ssh::Server, tcp::TcpHandler, - tcp_listener::get_tcp_listener, telemetry::{ TELEMETRY_COUNTER_NETWORK_RX, TELEMETRY_COUNTER_NETWORK_TX, TELEMETRY_COUNTER_SNI_CONNECTIONS, TELEMETRY_GAUGE_CPU_USAGE, TELEMETRY_GAUGE_TOTAL_MEMORY, diff --git a/src/lib.rs b/src/lib.rs index c714af0..f507bef 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -64,11 +64,9 @@ mod reactor; mod sock_addr_alias; mod ssh; mod tcp; -mod tcp_listener; mod telemetry; mod tls; mod udp; -mod udp_listener; // Data collected from the system and displayed on the admin interface. #[derive(Default, Clone)] diff --git a/src/tcp.rs b/src/tcp.rs index 35fbf92..5cac2a2 100644 --- a/src/tcp.rs +++ b/src/tcp.rs @@ -7,7 +7,6 @@ use crate::{ ip::IpFilter, reactor::TcpReactor, ssh::connection_handler::SshTunnelHandler, - tcp_listener::get_tcp_listener, telemetry::{TELEMETRY_COUNTER_TCP_CONNECTIONS, TELEMETRY_KEY_PORT}, }; use ahash::RandomState; @@ -15,6 +14,7 @@ use bon::Builder; use color_eyre::eyre::Context; use dashmap::DashMap; use metrics::counter; +use sandhole_socket::tcp_listener::get_tcp_listener; use tokio::{io::copy_bidirectional_with_sizes, time::timeout}; // Service that handles creating TCP sockets for reverse forwarding connections. diff --git a/src/udp.rs b/src/udp.rs index 6524b2a..01cc9e6 100644 --- a/src/udp.rs +++ b/src/udp.rs @@ -2,6 +2,7 @@ use std::{ collections::HashSet, mem::size_of, net::{IpAddr, SocketAddr}, + pin::pin, sync::Arc, time::Duration, }; @@ -14,33 +15,23 @@ use crate::{ reactor::UdpReactor, ssh::connection_handler::{SshChannel, SshTunnelHandler}, telemetry::{TELEMETRY_COUNTER_UDP_CONNECTIONS, TELEMETRY_KEY_PORT}, - udp_listener::get_udp_socket, }; use ahash::RandomState; use bon::Builder; use color_eyre::eyre::Context; use dashmap::DashMap; -use futures_util::pin_mut; use metrics::counter; +use sandhole_socket::{ + udp_listener::get_udp_socket, + udp_over_tcp::{datagram_buffer, deserialize_datagram, serialize_datagram}, +}; use tokio::{ - io::{AsyncReadExt, AsyncWriteExt, WriteHalf}, + io::{AsyncWriteExt, WriteHalf}, select, sync::Mutex, time::sleep, }; -pub const MAX_PACKET_SIZE: usize = size_of::() + u16::MAX as usize; - -// Creates and returns a buffer on the heap with enough space to contain any possible -// UDP datagram. -// -// This is put on the heap and in a separate function to avoid the 64k buffer from ending -// up on the stack and blowing up the size of the futures using it. -#[inline] -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. struct UdpSocketHandler { write: Arc>>, @@ -102,10 +93,7 @@ impl UdpPortHandler for Arc { // 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]); + serialize_datagram(&mut read_buf[..], &buf[..len]); // Check for an existing SSH channel if let Some(mut entry) = clone.sockets.get_mut(&(port, address)) { @@ -145,7 +133,7 @@ impl UdpPortHandler for Arc { } let udp_write = Arc::clone(&socket); - let (mut ssh_read, mut ssh_write) = tokio::io::split(channel); + let (ssh_read, mut ssh_write) = tokio::io::split(channel); if let Err(error) = ssh_write .write_all(&read_buf[..size_of::() + len]) @@ -159,37 +147,37 @@ impl UdpPortHandler for Arc { let clone_2 = Arc::clone(&clone); let _read_task = DroppableHandle(tokio::spawn(async move { let mut write_buf = datagram_buffer(); + let mut ssh_read = pin!(ssh_read); loop { let sleep_fut = sleep(clone_2.udp_timeout); let read_fut = async { - match ssh_read.read_u16().await { + match deserialize_datagram( + &mut write_buf[..], + &mut ssh_read, + ) + .await + { + Err(error) => { + #[cfg(not(coverage_nightly))] + tracing::warn!(%port, %error, "Error deserializing UDP datagram from SSH channel."); + false + } 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 UDP datagram from SSH channel."); - return false; - } else if let Err(error) = udp_write - .send_to(&write_buf[..len as usize], address) + if let Err(error) = udp_write + .send_to(&write_buf[..len], address) .await { #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error reading from SSH channel for UDP."); - return false; + tracing::warn!(%port, %error, "Error sending UDP datagram."); + false + } else { + true } - true - } - Err(error) => { - #[cfg(not(coverage_nightly))] - tracing::warn!(%port, %error, "Error reading UDP datagram size from SSH channel."); - false } } }; - pin_mut!(sleep_fut); - pin_mut!(read_fut); + let sleep_fut = pin!(sleep_fut); + let read_fut = pin!(read_fut); select! { success = read_fut => if !success { break; }, _ = sleep_fut => break, diff --git a/udp_over_tcp/Cargo.toml b/udp_over_tcp/Cargo.toml new file mode 100644 index 0000000..18f93ff --- /dev/null +++ b/udp_over_tcp/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "sandhole_udp_over_tcp" +version = "0.1.0" +edition = "2024" +rust-version = "1.88.0" +description = "Proxy UDP traffic for Sandhole via SSH." +repository = "https://github.com/EpicEric/sandhole" +homepage = "https://sandhole.com.br/" +license = "MIT" +authors = ["Eric Rodrigues Pires "] +readme = "README.md" +keywords = ["ssh", "proxy", "reverse-proxy", "tcp", "udp"] +categories = ["network-programming"] + +[dependencies] +ahash = { workspace = true } +clap = { workspace = true } +color-eyre = { workspace = true } +dashmap = { workspace = true } +sandhole_socket = { workspace = true } +tokio = { workspace = true } diff --git a/udp_over_tcp/src/main.rs b/udp_over_tcp/src/main.rs new file mode 100644 index 0000000..7adf31b --- /dev/null +++ b/udp_over_tcp/src/main.rs @@ -0,0 +1,261 @@ +use std::{ + future::Future, + net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, + num::NonZeroU16, + pin::{Pin, pin}, + sync::Arc, +}; + +use ahash::RandomState; +use clap::Parser; +use color_eyre::eyre::eyre; +use dashmap::DashMap; +use sandhole_socket::{ + tcp_listener::get_tcp_listener, + udp_listener::get_udp_socket, + udp_over_tcp::{datagram_buffer, deserialize_datagram, serialize_datagram}, +}; +use tokio::{io::AsyncWriteExt, net::TcpSocket, select, sync::mpsc}; + +#[doc(hidden)] +#[derive(Debug, Parser, PartialEq)] +#[command(version, about, long_about = None)] +struct Cli { + /// UDP address to proxy. + #[arg(long, value_name = "IP_ADDRESS")] + udp_address: Option, + + /// UDP port to proxy. + #[arg(long, value_name = "PORT")] + udp_port: NonZeroU16, + + /// TCP address to bind to. + #[arg(long, value_name = "IP_ADDRESS")] + tcp_address: Option, + + /// TCP port to bind to. + #[arg(long, value_name = "PORT")] + tcp_port: NonZeroU16, + + /// Start a UDP server for local forwarding. + #[arg(long)] + local_forwarding: bool, +} + +type FutType = Pin> + Send + Sync + 'static>>; + +// Starts a TCP server that redirects requests to the UDP socket. +async fn remote_forwarding(args: Cli) -> color_eyre::Result<()> { + let listener = get_tcp_listener(( + args.tcp_address + .unwrap_or(IpAddr::V6(Ipv6Addr::UNSPECIFIED)), + u16::from(args.tcp_port), + ))?; + let udp_address = ( + args.udp_address.unwrap_or(IpAddr::V4(Ipv4Addr::LOCALHOST)), + u16::from(args.udp_port), + ); + loop { + let (stream, _) = listener.accept().await?; + let udp_socket_read = Arc::new(get_udp_socket((IpAddr::V4(Ipv4Addr::LOCALHOST), 0))?); + udp_socket_read.connect(udp_address).await?; + let udp_socket_write = Arc::clone(&udp_socket_read); + tokio::spawn(async move { + let (read_stream, write_stream) = stream.into_split(); + + let read_fut: FutType = Box::pin(async move { + let mut read_stream = pin!(read_stream); + let mut buf = datagram_buffer(); + loop { + let len = deserialize_datagram(&mut buf[..], &mut read_stream).await?; + udp_socket_write.send(&buf[..len]).await?; + } + }); + + let write_fut: FutType = Box::pin(async move { + let mut write_stream = pin!(write_stream); + let mut datagram_buf = datagram_buffer(); + let mut buf = datagram_buffer(); + loop { + let datagram_len = udp_socket_read.recv(&mut datagram_buf[..]).await?; + let len = serialize_datagram(&mut buf[..], &datagram_buf[..datagram_len]); + write_stream.write_all(&buf[..len]).await?; + } + }); + + select! { + result = read_fut => if let Err(error) = result { + eprintln!("{error}"); + }, + result = write_fut => if let Err(error) = result { + eprintln!("{error}"); + }, + } + }); + } +} + +// Starts an UDP socket that redirects datagrams to the TCP client. +async fn local_forwarding(args: Cli) -> color_eyre::Result<()> { + let udp_socket_read = Arc::new(get_udp_socket(( + args.udp_address + .unwrap_or(IpAddr::V6(Ipv6Addr::UNSPECIFIED)), + u16::from(args.udp_port), + ))?); + + let read_fut: FutType = Box::pin(async move { + let mut datagram_buf = datagram_buffer(); + let mut buf = datagram_buffer(); + let conns: Arc>, RandomState>> = Arc::default(); + loop { + let (datagram_len, address) = udp_socket_read.recv_from(&mut datagram_buf[..]).await?; + let len = serialize_datagram(&mut buf[..], &datagram_buf[..datagram_len]); + + // Attempt to send data + let mut data = buf[..len].to_vec(); + if let Some(conn) = conns.get(&address) { + let conn = conn.clone(); + match conn.send(data).await { + Ok(_) => continue, + Err(error) => { + data = error.0; + conns.remove(&address); + } + } + } + + // No valid TCP connection; create a new one + let (serialized_datagram_tx, mut serialized_datagram_rx) = + mpsc::channel::>(128); + serialized_datagram_tx + .send(data) + .await + .expect("empty channel"); + conns.entry(address).insert_entry(serialized_datagram_tx); + let conns = Arc::clone(&conns); + let udp_socket_write = Arc::clone(&udp_socket_read); + tokio::spawn(async move { + let tcp_socket = if args.tcp_address.is_some_and(|address| address.is_ipv6()) { + TcpSocket::new_v6() + } else { + TcpSocket::new_v4() + } + .expect("should create TCP socket"); + tcp_socket + .set_nodelay(true) + .expect("should disable Nagle's algorithm"); + let tcp_stream = match tcp_socket + .connect(SocketAddr::new( + args.tcp_address.unwrap_or(IpAddr::V4(Ipv4Addr::LOCALHOST)), + u16::from(args.tcp_port), + )) + .await + { + Ok(tcp_stream) => tcp_stream, + Err(error) => { + eprintln!("Failed to connect to TCP address: {error}"); + conns.remove(&address); + return; + } + }; + let (read_stream, write_stream) = tcp_stream.into_split(); + + let mut write_fut: FutType = Box::pin(async move { + let mut write_stream = pin!(write_stream); + while let Some(data) = serialized_datagram_rx.recv().await { + write_stream.write_all(&data[..]).await?; + } + Ok(()) + }); + + let mut read_fut: FutType = Box::pin(async move { + let mut buf = datagram_buffer(); + let mut read_stream = pin!(read_stream); + loop { + let len = deserialize_datagram(&mut buf[..], &mut read_stream).await?; + udp_socket_write.send_to(&buf[..len], address).await?; + } + }); + + select! { + result = &mut write_fut => if let Err(error) = result { + eprintln!("Error writing to TCP client: {error}"); + }, + result = &mut read_fut => if let Err(error) = result { + eprintln!("Error writing to UDP socket: {error}"); + }, + } + + conns.remove(&address); + }); + } + }); + + read_fut.await +} + +#[tokio::main] +async fn main() -> color_eyre::Result<()> { + color_eyre::install()?; + let args = Cli::parse(); + + if args.local_forwarding { + select! { + result = local_forwarding(args) => result?, + signal = wait_for_signal() => { + return Err(eyre!("Received {signal}, terminating...")); + } + } + } else { + select! { + result = remote_forwarding(args) => result?, + signal = wait_for_signal() => { + return Err(eyre!("Received {signal}, terminating...")); + } + } + } + + Ok(()) +} + +#[cfg(unix)] +async fn wait_for_signal() -> &'static str { + use tokio::signal::unix::{SignalKind, signal}; + + let mut signal_terminate = signal(SignalKind::terminate()).expect("valid signal"); + let mut signal_interrupt = signal(SignalKind::interrupt()).expect("valid signal"); + + tokio::select! { + _ = signal_terminate.recv() => { + "SIGTERM" + }, + _ = signal_interrupt.recv() => { + "SIGINT" + }, + } +} + +#[cfg(windows)] +async fn wait_for_signal() -> &'static str { + use tokio::signal::windows; + + let mut signal_c = windows::ctrl_c().expect("valid signal"); + let mut signal_break = windows::ctrl_break().expect("valid signal"); + let mut signal_close = windows::ctrl_close().expect("valid signal"); + let mut signal_shutdown = windows::ctrl_shutdown().expect("valid signal"); + + tokio::select! { + _ = signal_c.recv() => { + "CTRL_C" + }, + _ = signal_break.recv() => { + "CTRL_BREAK" + }, + _ = signal_close.recv() => { + "CTRL_CLOSE" + }, + _ = signal_shutdown.recv() => { + "CTRL_SHUTDOWN" + }, + } +}