use anyhow::{anyhow, Context, Result};
use bytes::{Bytes, BytesMut};
use dashmap::DashMap;
use std::sync::Arc;
use std::time::Duration;
use quinn::{Connection, RecvStream, SendStream};
use std::net::SocketAddr;
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
use crate::common::tunnel::client_send_remote_request;
use super::remote::RemoteRequest;
const MAX_DATAGRAM: usize = 65_535;
const SESSION_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
const SESSION_CHANNEL_CAPACITY: usize = 4096;
const UDP_RECV_POOL_BYTES: usize = MAX_DATAGRAM * 8;
pub(crate) async fn write_datagram(send: &mut SendStream, payload: &[u8]) -> Result<()> {
let len = u16::try_from(payload.len())
.map_err(|_| anyhow!("UDP datagram of {} bytes exceeds 65535", payload.len()))?;
send.write_all(&len.to_le_bytes()).await?;
send.write_all(payload).await?;
Ok(())
}
pub(crate) async fn read_datagram<'a>(
recv: &mut RecvStream,
buf: &'a mut [u8],
) -> Result<&'a [u8]> {
let mut len_buf = [0u8; 2];
recv.read_exact(&mut len_buf).await?;
let len = u16::from_le_bytes(len_buf) as usize;
if len > buf.len() {
return Err(anyhow!(
"datagram length {} exceeds local buffer {}",
len,
buf.len()
));
}
recv.read_exact(&mut buf[..len]).await?;
Ok(&buf[..len])
}
pub async fn tunnel_udp_stream(
udp_socket: Arc<UdpSocket>,
udp_address: SocketAddr,
send_channel: SendStream,
recv_channel: RecvStream,
) -> Result<()> {
let (l2q, q2l) = tokio::join!(
pump_socket_to_stream(udp_socket.clone(), udp_address, send_channel),
pump_stream_to_socket(udp_socket, udp_address, recv_channel),
);
if let Err(e) = l2q {
debug!("local→quic error: {}", e);
}
if let Err(e) = q2l {
debug!("quic→local error: {}", e);
}
Ok(())
}
async fn pump_socket_to_stream(
udp_socket: Arc<UdpSocket>,
udp_address: SocketAddr,
mut send: SendStream,
) -> Result<()> {
let mut buf = vec![0u8; MAX_DATAGRAM];
loop {
let (len, received_addr) = udp_socket.recv_from(&mut buf).await?;
if received_addr != udp_address {
debug!(peer = %received_addr, expected = %udp_address, "dropping datagram from unexpected source");
continue;
}
write_datagram(&mut send, &buf[..len]).await?;
}
}
async fn pump_stream_to_socket(
udp_socket: Arc<UdpSocket>,
udp_address: SocketAddr,
mut recv: RecvStream,
) -> Result<()> {
let mut buf = vec![0u8; MAX_DATAGRAM];
loop {
let payload = read_datagram(&mut recv, &mut buf).await?;
udp_socket.send_to(payload, &udp_address).await?;
}
}
pub async fn tunnel_udp_client(quic_connection: Connection, remote: RemoteRequest) -> Result<()> {
let listen_addr = remote.local_socket_addr();
let udp_socket = Arc::new(UdpSocket::bind(listen_addr).await?);
info!("listening on {}", listen_addr);
let sessions: Arc<DashMap<SocketAddr, mpsc::Sender<Bytes>>> = Arc::new(DashMap::new());
let mut recv_buf = BytesMut::with_capacity(UDP_RECV_POOL_BYTES);
loop {
if recv_buf.capacity() < MAX_DATAGRAM {
recv_buf.reserve(UDP_RECV_POOL_BYTES);
}
recv_buf.resize(MAX_DATAGRAM, 0);
let (n, src) = udp_socket.recv_from(&mut recv_buf[..]).await?;
let payload = recv_buf.split_to(n).freeze();
let mut existing = sessions.get(&src).map(|e| e.value().clone());
if let Some(tx) = &existing {
if tx.is_closed() {
sessions.remove(&src);
existing = None;
}
}
let sender = match existing {
Some(tx) => tx,
None => {
let (tx, rx) = mpsc::channel(SESSION_CHANNEL_CAPACITY);
sessions.insert(src, tx.clone());
spawn_udp_session(
quic_connection.clone(),
remote.clone(),
udp_socket.clone(),
src,
rx,
sessions.clone(),
);
tx
}
};
if let Err(e) = sender.try_send(payload) {
debug!(peer = %src, "dropping UDP datagram: {}", e);
}
}
}
fn spawn_udp_session(
quic_connection: Connection,
remote: RemoteRequest,
udp_socket: Arc<UdpSocket>,
source: SocketAddr,
rx: mpsc::Receiver<Bytes>,
sessions: Arc<DashMap<SocketAddr, mpsc::Sender<Bytes>>>,
) {
tokio::spawn(async move {
debug!(peer = %source, "opening UDP session");
if let Err(e) = run_udp_session(quic_connection, remote, udp_socket, source, rx).await {
warn!(peer = %source, "UDP session ended: {}", e);
}
sessions.remove(&source);
debug!(peer = %source, "UDP session removed");
});
}
async fn run_udp_session(
quic_connection: Connection,
remote: RemoteRequest,
udp_socket: Arc<UdpSocket>,
source: SocketAddr,
mut rx: mpsc::Receiver<Bytes>,
) -> Result<()> {
let (mut send_channel, mut recv_channel) = quic_connection.open_bi().await?;
client_send_remote_request(&remote, &mut send_channel, &mut recv_channel).await?;
let local_to_quic = async {
loop {
match tokio::time::timeout(SESSION_IDLE_TIMEOUT, rx.recv()).await {
Ok(Some(payload)) => write_datagram(&mut send_channel, &payload).await?,
Ok(None) => return Ok::<(), anyhow::Error>(()), Err(_) => {
debug!(peer = %source, "UDP session idle timeout");
return Ok(());
}
}
}
};
let quic_to_local = async {
let mut buf = vec![0u8; MAX_DATAGRAM];
loop {
let payload = read_datagram(&mut recv_channel, &mut buf).await?;
udp_socket.send_to(payload, &source).await?;
}
};
tokio::select! {
r = local_to_quic => r,
r = quic_to_local => r,
}
}
pub async fn tunnel_udp_server(
recv_channel: RecvStream,
send_channel: SendStream,
request: RemoteRequest,
) -> Result<()> {
let remote_addr: SocketAddr = request
.remote_addr_string()
.ok_or_else(|| anyhow!("UDP server tunnel requires a host:port remote"))?
.parse()
.context("Failed to parse remote address")?;
let bind_addr = match remote_addr {
SocketAddr::V4(_) => "0.0.0.0:0",
SocketAddr::V6(_) => "[::]:0",
};
let udp_socket = Arc::new(UdpSocket::bind(bind_addr).await?);
debug!("connecting to {}", remote_addr);
tunnel_udp_stream(udp_socket, remote_addr, send_channel, recv_channel).await?;
Ok(())
}