use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use dashmap::DashMap;
use quinn::{Connection, RecvStream, SendStream};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream, UdpSocket};
use tokio::sync::mpsc;
use tracing::{debug, error, info, info_span, warn, Instrument};
use super::remote::{DynamicTarget, HostPort, OpenConn, RemoteRequest};
use super::tcp::{tunnel_tcp_stream, Counters, TunnelHandleOpt};
use super::tunnel::send_open_conn;
use super::udp::{read_datagram, write_datagram};
use anyhow::{anyhow, Result};
const SOCKS_MAX_DATAGRAM: usize = 65_535;
const SOCKS_UDP_CHANNEL_CAPACITY: usize = 4096;
const SOCKS_UDP_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
pub async fn tunnel_socks_client(
quic_connection: Connection,
remote: RemoteRequest,
handle: TunnelHandleOpt,
tunnel_id: u64,
) -> Result<()> {
let local_addr = remote.local_socket_addr();
let listener = TcpListener::bind(local_addr).await?;
info!(addr = %local_addr, proto = "socks5", "listening");
let local_counter = AtomicUsize::new(0);
loop {
let (mut local_conn, peer) = listener.accept().await?;
let connection = quic_connection.clone();
let remote = remote.clone();
let tunnel_handle = handle.clone();
let local_id = (local_counter.fetch_add(1, Ordering::Relaxed) + 1) as u64;
tokio::spawn(async move {
let request = match socks_handshake(&mut local_conn).await {
Ok(r) => r,
Err(e) => {
let span = info_span!("socks5", peer = %peer);
let _g = span.enter();
warn!(error = %e, "handshake failed");
return;
}
};
match request {
SocksRequest::Connect(target) => {
let (send, recv) = match connection.open_bi().await {
Ok(stream) => stream,
Err(e) => {
let span = info_span!("socks5", peer = %peer);
let _g = span.enter();
error!(error = %e, "failed to open quic stream");
return;
}
};
let conn_guard = tunnel_handle
.as_ref()
.map(|h| h.open_conn(Some(format!("{peer}=>{target}"))));
let conn_id = conn_guard.as_ref().map(|g| g.id()).unwrap_or(local_id);
let counters = conn_guard.as_ref().map(|g| g.counters());
let span = info_span!(
"conn",
conn_id,
tunnel_id,
peer = %peer,
target = %target,
proto = "socks5/tcp",
);
async move {
info!("conn opened");
let started = std::time::Instant::now();
let result = start_client_dynamic_tunnel(
local_conn,
send,
recv,
tunnel_id,
target,
counters.clone(),
)
.await;
let dur_ms = started.elapsed().as_millis() as u64;
let snap = counters.as_ref().map(|c| c.snapshot());
match (&result, snap) {
(Ok(()), Some((bytes_in, bytes_out))) => {
info!(bytes_in, bytes_out, dur_ms, "conn closed")
}
(Ok(()), None) => info!(dur_ms, "conn closed"),
(Err(e), _) => warn!(dur_ms, error = %e, "conn closed (error)"),
}
drop(conn_guard);
}
.instrument(span)
.await;
}
SocksRequest::UdpAssociate => {
let span = info_span!("socks5-udp", peer = %peer);
async move {
debug!("UDP ASSOCIATE requested");
if let Err(e) = handle_socks_udp_associate(
connection,
local_conn,
&remote,
tunnel_handle,
peer,
tunnel_id,
)
.await
{
error!(error = %e, "UDP ASSOCIATE failed");
}
}
.instrument(span)
.await;
}
}
});
}
}
async fn start_client_dynamic_tunnel(
mut socks_conn: TcpStream,
mut send_channel: SendStream,
mut recv_channel: RecvStream,
tunnel_id: u64,
target: HostPort,
counters: Counters,
) -> Result<()> {
send_open_conn(
&OpenConn {
tunnel_id,
dynamic: Some(DynamicTarget::Tcp(target)),
},
&mut send_channel,
&mut recv_channel,
)
.await?;
socks_conn
.write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0])
.await?;
tunnel_tcp_stream(socks_conn, send_channel, recv_channel, counters).await?;
Ok(())
}
enum SocksRequest {
Connect(HostPort),
UdpAssociate,
}
async fn socks_handshake(conn: &mut TcpStream) -> Result<SocksRequest> {
let mut buf = [0u8; 256];
conn.read_exact(&mut buf[..2]).await?;
if buf[0] != 0x05 {
return Err(anyhow!("Unsupported SOCKS version: {}", buf[0]));
}
let methods_len = buf[1] as usize;
conn.read_exact(&mut buf[..methods_len]).await?;
conn.write_all(&[0x05, 0x00]).await?;
conn.read_exact(&mut buf[..4]).await?;
let cmd = buf[1];
let atyp = buf[3];
let target = read_socks_addr(conn, atyp).await?;
match cmd {
0x01 => Ok(SocksRequest::Connect(target)),
0x03 => Ok(SocksRequest::UdpAssociate),
_ => {
conn.write_all(&[0x05, 0x07]).await?;
Err(anyhow!("Unsupported SOCKS command: {}", cmd))
}
}
}
async fn read_socks_addr(conn: &mut TcpStream, atyp: u8) -> Result<HostPort> {
match atyp {
0x01 => {
let mut addr = [0u8; 4];
conn.read_exact(&mut addr).await?;
let mut port = [0u8; 2];
conn.read_exact(&mut port).await?;
Ok(HostPort::new(
Ipv4Addr::from(addr).to_string(),
u16::from_be_bytes(port),
))
}
0x03 => {
let mut len = [0u8; 1];
conn.read_exact(&mut len).await?;
let mut domain = vec![0u8; len[0] as usize];
conn.read_exact(&mut domain).await?;
let mut port = [0u8; 2];
conn.read_exact(&mut port).await?;
Ok(HostPort::new(
String::from_utf8_lossy(&domain).into_owned(),
u16::from_be_bytes(port),
))
}
0x04 => {
let mut addr = [0u8; 16];
conn.read_exact(&mut addr).await?;
let mut port = [0u8; 2];
conn.read_exact(&mut port).await?;
Ok(HostPort::new(
Ipv6Addr::from(addr).to_string(),
u16::from_be_bytes(port),
))
}
other => {
conn.write_all(&[0x05, 0x08]).await?;
Err(anyhow!("Unsupported address type: {}", other))
}
}
}
async fn handle_socks_udp_associate(
quic_connection: Connection,
mut tcp_conn: TcpStream,
original_remote: &RemoteRequest,
handle: TunnelHandleOpt,
socks_peer: SocketAddr,
tunnel_id: u64,
) -> Result<()> {
let listen_ip = original_remote.local_socket_addr().ip();
let bind_addr = SocketAddr::new(listen_ip, 0);
let udp_socket = Arc::new(UdpSocket::bind(bind_addr).await?);
let bound = udp_socket.local_addr()?;
debug!(addr = %bound, "UDP ASSOCIATE relay bound");
write_socks_udp_associate_reply(&mut tcp_conn, bound).await?;
let relay_handle = tokio::spawn({
let connection = quic_connection.clone();
let socket = udp_socket.clone();
let handle = handle.clone();
async move {
if let Err(e) =
run_socks_udp_relay(connection, tunnel_id, socket, handle, socks_peer).await
{
debug!(error = %e, "UDP relay ended");
}
}
});
let mut buf = [0u8; 64];
while let Ok(n) = tcp_conn.read(&mut buf).await {
if n == 0 {
break;
}
}
relay_handle.abort();
debug!("UDP ASSOCIATE control closed");
Ok(())
}
async fn write_socks_udp_associate_reply(conn: &mut TcpStream, addr: SocketAddr) -> Result<()> {
let mut reply = vec![0x05, 0x00, 0x00];
match addr.ip() {
IpAddr::V4(v4) => {
reply.push(0x01);
reply.extend_from_slice(&v4.octets());
}
IpAddr::V6(v6) => {
reply.push(0x04);
reply.extend_from_slice(&v6.octets());
}
}
reply.extend_from_slice(&addr.port().to_be_bytes());
conn.write_all(&reply).await?;
Ok(())
}
async fn run_socks_udp_relay(
quic_connection: Connection,
tunnel_id: u64,
udp_socket: Arc<UdpSocket>,
handle: TunnelHandleOpt,
socks_peer: SocketAddr,
) -> Result<()> {
let conns: Arc<DashMap<(SocketAddr, HostPort), mpsc::Sender<Bytes>>> = Arc::new(DashMap::new());
let mut buf = vec![0u8; SOCKS_MAX_DATAGRAM];
loop {
let (n, src) = udp_socket.recv_from(&mut buf).await?;
let (target, payload) = match parse_socks_udp_header(&buf[..n]) {
Ok(v) => v,
Err(e) => {
debug!(peer = %src, error = %e, "invalid UDP datagram");
continue;
}
};
let key = (src, target.clone());
let mut existing = conns.get(&key).map(|e| e.value().clone());
if let Some(tx) = &existing {
if tx.is_closed() {
conns.remove(&key);
existing = None;
}
}
let tx = match existing {
Some(tx) => tx,
None => {
let (tx, rx) = mpsc::channel(SOCKS_UDP_CHANNEL_CAPACITY);
conns.insert(key.clone(), tx.clone());
spawn_socks_udp_conn(
quic_connection.clone(),
tunnel_id,
udp_socket.clone(),
src,
target.clone(),
rx,
conns.clone(),
handle.clone(),
socks_peer,
);
tx
}
};
if let Err(e) = tx.try_send(Bytes::copy_from_slice(payload)) {
debug!(peer = %src, target = %target, error = %e, "dropping udp datagram");
}
}
}
#[allow(clippy::too_many_arguments)]
fn spawn_socks_udp_conn(
quic_connection: Connection,
tunnel_id: u64,
udp_socket: Arc<UdpSocket>,
source: SocketAddr,
target: HostPort,
rx: mpsc::Receiver<Bytes>,
conns: Arc<DashMap<(SocketAddr, HostPort), mpsc::Sender<Bytes>>>,
handle: TunnelHandleOpt,
socks_peer: SocketAddr,
) {
let key = (source, target.clone());
tokio::spawn(async move {
let conn_guard = handle
.as_ref()
.map(|h| h.open_conn(Some(format!("{socks_peer}=>{target}"))));
let conn_id = conn_guard.as_ref().map(|g| g.id()).unwrap_or(0);
let counters = conn_guard.as_ref().map(|g| g.counters());
let span = info_span!(
"conn",
conn_id,
tunnel_id,
peer = %socks_peer,
target = %target,
proto = "socks5/udp",
);
async move {
info!("conn opened");
let started = std::time::Instant::now();
let result = run_socks_udp_conn(
quic_connection,
tunnel_id,
udp_socket,
source,
target,
rx,
counters.clone(),
)
.await;
let dur_ms = started.elapsed().as_millis() as u64;
let snap = counters.as_ref().map(|c| c.snapshot());
match (&result, snap) {
(Ok(()), Some((bytes_in, bytes_out))) => {
info!(bytes_in, bytes_out, dur_ms, "conn closed")
}
(Ok(()), None) => info!(dur_ms, "conn closed"),
(Err(e), _) => warn!(dur_ms, error = %e, "conn closed (error)"),
}
}
.instrument(span)
.await;
drop(conn_guard);
conns.remove(&key);
});
}
#[allow(clippy::too_many_arguments)]
async fn run_socks_udp_conn(
quic_connection: Connection,
tunnel_id: u64,
udp_socket: Arc<UdpSocket>,
source: SocketAddr,
target: HostPort,
mut rx: mpsc::Receiver<Bytes>,
counters: Counters,
) -> Result<()> {
let (mut send_channel, mut recv_channel) = quic_connection.open_bi().await?;
send_open_conn(
&OpenConn {
tunnel_id,
dynamic: Some(DynamicTarget::Udp(target.clone())),
},
&mut send_channel,
&mut recv_channel,
)
.await?;
let local_to_quic = async {
loop {
match tokio::time::timeout(SOCKS_UDP_IDLE_TIMEOUT, rx.recv()).await {
Ok(Some(payload)) => {
write_datagram(&mut send_channel, &payload).await?;
if let Some(c) = counters.as_ref() {
c.add_out(payload.len() as u64);
}
}
Ok(None) => return Ok::<(), anyhow::Error>(()),
Err(_) => {
debug!("idle timeout");
return Ok(());
}
}
}
};
let quic_to_local = async {
let mut buf = vec![0u8; SOCKS_MAX_DATAGRAM];
let mut wrap = Vec::with_capacity(SOCKS_MAX_DATAGRAM + 32);
loop {
let payload = read_datagram(&mut recv_channel, &mut buf).await?;
wrap_socks_udp_reply(&target, payload, &mut wrap);
udp_socket.send_to(&wrap, source).await?;
if let Some(c) = counters.as_ref() {
c.add_in(payload.len() as u64);
}
}
};
tokio::select! {
r = local_to_quic => r,
r = quic_to_local => r,
}
}
fn parse_socks_udp_header(buf: &[u8]) -> Result<(HostPort, &[u8])> {
if buf.len() < 4 {
return Err(anyhow!("UDP header too short"));
}
if buf[2] != 0 {
return Err(anyhow!("UDP fragmentation not supported"));
}
let atyp = buf[3];
let (target, hdr_len) = match atyp {
0x01 => {
if buf.len() < 4 + 4 + 2 {
return Err(anyhow!("truncated IPv4 header"));
}
let ip = Ipv4Addr::new(buf[4], buf[5], buf[6], buf[7]);
let port = u16::from_be_bytes([buf[8], buf[9]]);
(HostPort::new(ip.to_string(), port), 10)
}
0x04 => {
if buf.len() < 4 + 16 + 2 {
return Err(anyhow!("truncated IPv6 header"));
}
let mut octets = [0u8; 16];
octets.copy_from_slice(&buf[4..20]);
let ip = Ipv6Addr::from(octets);
let port = u16::from_be_bytes([buf[20], buf[21]]);
(HostPort::new(ip.to_string(), port), 22)
}
0x03 => {
if buf.len() < 5 {
return Err(anyhow!("truncated domain header"));
}
let len = buf[4] as usize;
if buf.len() < 5 + len + 2 {
return Err(anyhow!("truncated domain header"));
}
let domain = String::from_utf8_lossy(&buf[5..5 + len]).into_owned();
let port = u16::from_be_bytes([buf[5 + len], buf[5 + len + 1]]);
(HostPort::new(domain, port), 5 + len + 2)
}
other => return Err(anyhow!("unknown ATYP: {}", other)),
};
Ok((target, &buf[hdr_len..]))
}
fn wrap_socks_udp_reply(target: &HostPort, payload: &[u8], buf: &mut Vec<u8>) {
buf.clear();
buf.extend_from_slice(&[0, 0, 0]);
if let Ok(ip) = target.host.parse::<Ipv4Addr>() {
buf.push(0x01);
buf.extend_from_slice(&ip.octets());
} else if let Ok(ip) = target.host.parse::<Ipv6Addr>() {
buf.push(0x04);
buf.extend_from_slice(&ip.octets());
} else {
buf.push(0x03);
let bytes = target.host.as_bytes();
let len = bytes.len().min(u8::MAX as usize);
buf.push(len as u8);
buf.extend_from_slice(&bytes[..len]);
}
buf.extend_from_slice(&target.port.to_be_bytes());
buf.extend_from_slice(payload);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_udp_header_ipv4() {
let buf = [0, 0, 0, 1, 1, 2, 3, 4, 0, 53, b'a', b'b', b'c'];
let (target, payload) = parse_socks_udp_header(&buf).unwrap();
assert_eq!(target.host, "1.2.3.4");
assert_eq!(target.port, 53);
assert_eq!(payload, b"abc");
}
#[test]
fn parse_udp_header_domain() {
let mut buf = vec![0, 0, 0, 3, 11];
buf.extend_from_slice(b"example.com");
buf.extend_from_slice(&80u16.to_be_bytes());
buf.extend_from_slice(b"hi");
let (target, payload) = parse_socks_udp_header(&buf).unwrap();
assert_eq!(target.host, "example.com");
assert_eq!(target.port, 80);
assert_eq!(payload, b"hi");
}
#[test]
fn parse_udp_header_ipv6() {
let mut buf = vec![0, 0, 0, 4];
buf.extend_from_slice(&Ipv6Addr::LOCALHOST.octets());
buf.extend_from_slice(&443u16.to_be_bytes());
buf.extend_from_slice(b"x");
let (target, payload) = parse_socks_udp_header(&buf).unwrap();
assert_eq!(target.host, "::1");
assert_eq!(target.port, 443);
assert_eq!(payload, b"x");
}
#[test]
fn parse_udp_header_rejects_fragmentation() {
let buf = [0, 0, 1, 1, 1, 2, 3, 4, 0, 53];
assert!(parse_socks_udp_header(&buf).is_err());
}
#[test]
fn wrap_reply_ipv4_roundtrip() {
let target = HostPort::new("8.8.8.8", 53);
let mut out = Vec::new();
wrap_socks_udp_reply(&target, b"hello", &mut out);
let (parsed, payload) = parse_socks_udp_header(&out).unwrap();
assert_eq!(parsed, target);
assert_eq!(payload, b"hello");
}
#[test]
fn wrap_reply_domain_roundtrip() {
let target = HostPort::new("example.com", 80);
let mut out = Vec::new();
wrap_socks_udp_reply(&target, b"data", &mut out);
let (parsed, payload) = parse_socks_udp_header(&out).unwrap();
assert_eq!(parsed, target);
assert_eq!(payload, b"data");
}
}