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, info_span, warn, Instrument};
use crate::common::remote::OpenConn;
use crate::common::tcp::{Counters, TunnelHandleOpt};
use crate::common::tunnel::send_open_conn;
use super::remote::RemoteRequest;
const MAX_DATAGRAM: usize = 65_535;
const CONN_IDLE_TIMEOUT: Duration = Duration::from_secs(60);
const CONN_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,
counters: Counters,
) -> Result<()> {
let (l2q, q2l) = tokio::join!(
pump_socket_to_stream(
udp_socket.clone(),
udp_address,
send_channel,
counters.clone()
),
pump_stream_to_socket(udp_socket, udp_address, recv_channel, counters),
);
if let Err(e) = l2q {
debug!(direction = "tx", error = %e, "udp pump error");
}
if let Err(e) = q2l {
debug!(direction = "rx", error = %e, "udp pump error");
}
Ok(())
}
async fn pump_socket_to_stream(
udp_socket: Arc<UdpSocket>,
udp_address: SocketAddr,
mut send: SendStream,
counters: Counters,
) -> 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!(from = %received_addr, expected = %udp_address, "dropping unexpected udp source");
continue;
}
write_datagram(&mut send, &buf[..len]).await?;
if let Some(c) = counters.as_ref() {
c.add_out(len as u64);
}
}
}
async fn pump_stream_to_socket(
udp_socket: Arc<UdpSocket>,
udp_address: SocketAddr,
mut recv: RecvStream,
counters: Counters,
) -> 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?;
if let Some(c) = counters.as_ref() {
c.add_in(payload.len() as u64);
}
}
}
pub async fn tunnel_udp_client(
quic_connection: Connection,
remote: RemoteRequest,
handle: TunnelHandleOpt,
tunnel_id: u64,
) -> Result<()> {
let listen_addr = remote.local_socket_addr();
let udp_socket = Arc::new(UdpSocket::bind(listen_addr).await?);
info!(addr = %listen_addr, proto = "udp", "listening");
let conns: 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 = conns.get(&src).map(|e| e.value().clone());
if let Some(tx) = &existing {
if tx.is_closed() {
conns.remove(&src);
existing = None;
}
}
let sender = match existing {
Some(tx) => tx,
None => {
let (tx, rx) = mpsc::channel(CONN_CHANNEL_CAPACITY);
conns.insert(src, tx.clone());
spawn_udp_conn(
quic_connection.clone(),
udp_socket.clone(),
src,
rx,
conns.clone(),
handle.clone(),
tunnel_id,
);
tx
}
};
if let Err(e) = sender.try_send(payload) {
debug!(peer = %src, error = %e, "dropping udp datagram");
}
}
}
#[allow(clippy::too_many_arguments)]
fn spawn_udp_conn(
quic_connection: Connection,
udp_socket: Arc<UdpSocket>,
source: SocketAddr,
rx: mpsc::Receiver<Bytes>,
conns: Arc<DashMap<SocketAddr, mpsc::Sender<Bytes>>>,
handle: TunnelHandleOpt,
tunnel_id: u64,
) {
tokio::spawn(async move {
let conn_guard = handle
.as_ref()
.map(|h| h.open_conn(Some(source.to_string())));
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 = %source, proto = "udp");
async move {
info!("conn opened");
let started = std::time::Instant::now();
let result = run_udp_conn(
quic_connection,
tunnel_id,
udp_socket,
source,
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), Some((bytes_in, bytes_out))) => {
warn!(bytes_in, bytes_out, dur_ms, error = %e, "conn closed (error)")
}
(Err(e), None) => warn!(dur_ms, error = %e, "conn closed (error)"),
}
}
.instrument(span)
.await;
drop(conn_guard);
conns.remove(&source);
});
}
async fn run_udp_conn(
quic_connection: Connection,
tunnel_id: u64,
udp_socket: Arc<UdpSocket>,
source: SocketAddr,
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: None,
},
&mut send_channel,
&mut recv_channel,
)
.await?;
let local_to_quic = async {
loop {
match tokio::time::timeout(CONN_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; MAX_DATAGRAM];
loop {
let payload = read_datagram(&mut recv_channel, &mut buf).await?;
udp_socket.send_to(payload, &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,
}
}
pub async fn tunnel_udp_server(
recv_channel: RecvStream,
send_channel: SendStream,
request: RemoteRequest,
counters: Counters,
) -> 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!(target = %remote_addr, "udp dial");
tunnel_udp_stream(
udp_socket,
remote_addr,
send_channel,
recv_channel,
counters,
)
.await?;
Ok(())
}