use anyhow::{anyhow, Context, Result};
use std::collections::HashMap;
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, Mutex};
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;
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(())
}
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,
mut send_channel: SendStream,
mut recv_channel: RecvStream,
) -> Result<()> {
let socket_for_recv = udp_socket.clone();
let local_to_quic = async move {
let mut buf = vec![0u8; MAX_DATAGRAM];
loop {
let (len, received_addr) = socket_for_recv.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_channel, &buf[..len]).await?;
}
#[allow(unreachable_code)]
Ok::<(), anyhow::Error>(())
};
let quic_to_local = async move {
let mut buf = vec![0u8; MAX_DATAGRAM];
loop {
let payload = read_datagram(&mut recv_channel, &mut buf).await?;
udp_socket.send_to(payload, &udp_address).await?;
}
#[allow(unreachable_code)]
Ok::<(), anyhow::Error>(())
};
let (l2q, q2l) = tokio::join!(local_to_quic, quic_to_local);
if let Err(e) = l2q {
debug!("local→quic error: {}", e);
}
if let Err(e) = q2l {
debug!("quic→local error: {}", e);
}
Ok(())
}
pub async fn tunnel_udp_client(quic_connection: Connection, remote: RemoteRequest) -> Result<()> {
let listen_addr = format!("{}:{}", remote.local_host, remote.local_port);
let udp_socket = Arc::new(UdpSocket::bind(&listen_addr).await?);
info!("listening on {}", listen_addr);
let sessions: Arc<Mutex<HashMap<SocketAddr, mpsc::Sender<Vec<u8>>>>> =
Arc::new(Mutex::new(HashMap::new()));
let mut buf = vec![0u8; MAX_DATAGRAM];
loop {
let (n, src) = udp_socket.recv_from(&mut buf).await?;
let payload = buf[..n].to_vec();
let sender = {
let mut map = sessions.lock().await;
if let Some(tx) = map.get(&src) {
if tx.is_closed() {
map.remove(&src);
}
}
if let Some(tx) = map.get(&src) {
tx.clone()
} else {
let (tx, rx) = mpsc::channel(SESSION_CHANNEL_CAPACITY);
map.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<Vec<u8>>,
sessions: Arc<Mutex<HashMap<SocketAddr, mpsc::Sender<Vec<u8>>>>>,
) {
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.lock().await.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<Vec<u8>>,
) -> 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(()), Err(_) => {
debug!(peer = %source, "UDP session idle timeout");
return Ok(());
}
}
}
#[allow(unreachable_code)]
Ok::<(), anyhow::Error>(())
};
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?;
}
#[allow(unreachable_code)]
Ok::<(), anyhow::Error>(())
};
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 = format!("{}:{}", request.remote_host, request.remote_port)
.parse()
.context("Failed to parse remote address")?;
let udp_socket = Arc::new(UdpSocket::bind("0.0.0.0:0").await?);
debug!("connecting to {}", remote_addr);
tunnel_udp_stream(udp_socket, remote_addr, send_channel, recv_channel).await?;
Ok(())
}