use std::collections::HashMap;
use std::io;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use tokio::net::{UdpSocket, lookup_host};
use crate::io::ConnLog;
use crate::socks;
use crate::stats::{EndpointStats, Registry};
const MAX_DATAGRAM: usize = 65_535;
#[derive(Debug)]
pub struct Relay {
client_side: UdpSocket,
upstream: UdpSocket,
registry: Arc<Registry>,
log: Option<Arc<ConnLog>>,
endpoints: HashMap<String, Arc<EndpointStats>>,
resolved: HashMap<(String, u16), SocketAddr>,
names: HashMap<IpAddr, String>,
client_addr: Option<SocketAddr>,
}
impl Relay {
pub async fn bind(registry: Arc<Registry>, log: Option<Arc<ConnLog>>) -> io::Result<Self> {
Ok(Self {
client_side: UdpSocket::bind(("127.0.0.1", 0)).await?,
upstream: UdpSocket::bind(("0.0.0.0", 0)).await?,
registry,
log,
endpoints: HashMap::new(),
resolved: HashMap::new(),
names: HashMap::new(),
client_addr: None,
})
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.client_side.local_addr()
}
pub async fn run(mut self) {
let mut from_client = vec![0u8; MAX_DATAGRAM];
let mut from_upstream = vec![0u8; MAX_DATAGRAM];
let mut scratch = Vec::with_capacity(MAX_DATAGRAM);
loop {
let event = tokio::select! {
result = self.client_side.recv_from(&mut from_client) => match result {
Ok((len, addr)) => Event::FromClient { len, addr },
Err(_) => return,
},
result = self.upstream.recv_from(&mut from_upstream) => match result {
Ok((len, addr)) => Event::FromUpstream { len, addr },
Err(_) => return,
},
};
match event {
Event::FromClient { len, addr } => {
self.client_addr = Some(addr);
self.forward_out(&from_client[..len]).await;
}
Event::FromUpstream { len, addr } => {
self.forward_back(addr, &from_upstream[..len], &mut scratch)
.await;
}
}
}
}
async fn forward_out(&mut self, raw: &[u8]) {
let datagram = match socks::parse_datagram(raw) {
Ok(datagram) => datagram,
Err(err) => {
self.note(format!("UDP -- dropped datagram: {err}"));
return;
}
};
let Some(target) = self.resolve(&datagram.host, datagram.port).await else {
return;
};
let stats = self.endpoint(&datagram.host);
match self.upstream.send_to(datagram.payload, target).await {
Ok(sent) => {
stats.add_egress(sent as u64);
stats.observe_ip(target.ip());
self.names.insert(target.ip(), datagram.host.clone());
if let Some(log) = &self.log {
log.event(format!(
"UDP -> {sent} bytes {}:{}",
datagram.host, datagram.port
));
}
}
Err(err) => self.note(format!("UDP -- send to {target} failed: {err}")),
}
}
async fn forward_back(&mut self, from: SocketAddr, payload: &[u8], scratch: &mut Vec<u8>) {
let Some(client_addr) = self.client_addr else {
return;
};
let host = self
.names
.get(&from.ip())
.cloned()
.unwrap_or_else(|| from.ip().to_string());
let stats = self.endpoint(&host);
socks::encode_datagram(from, payload, scratch);
match self.client_side.send_to(scratch, client_addr).await {
Ok(_) => {
stats.add_ingress(payload.len() as u64);
if let Some(log) = &self.log {
log.event(format!("UDP <- {} bytes {host}", payload.len()));
}
}
Err(err) => self.note(format!("UDP -- reply to client failed: {err}")),
}
}
fn endpoint(&mut self, host: &str) -> Arc<EndpointStats> {
if let Some(stats) = self.endpoints.get(host) {
return Arc::clone(stats);
}
let stats = self.registry.endpoint(host);
stats.add_connection();
self.endpoints.insert(host.to_owned(), Arc::clone(&stats));
stats
}
async fn resolve(&mut self, host: &str, port: u16) -> Option<SocketAddr> {
let key = (host.to_owned(), port);
if let Some(addr) = self.resolved.get(&key) {
return Some(*addr);
}
let addr = match lookup_host((host, port)).await {
Ok(mut addrs) => addrs.next(),
Err(err) => {
self.note(format!("UDP -- cannot resolve {host}:{port}: {err}"));
None
}
}?;
self.resolved.insert(key, addr);
Some(addr)
}
fn note(&self, message: String) {
if let Some(log) = &self.log {
log.event(message);
}
}
}
#[derive(Debug, Clone, Copy)]
enum Event {
FromClient { len: usize, addr: SocketAddr },
FromUpstream { len: usize, addr: SocketAddr },
}
#[cfg(test)]
mod tests {
use super::*;
async fn udp_echo() -> SocketAddr {
let socket = UdpSocket::bind(("127.0.0.1", 0)).await.unwrap();
let addr = socket.local_addr().unwrap();
tokio::spawn(async move {
let mut buf = vec![0u8; 2048];
while let Ok((len, from)) = socket.recv_from(&mut buf).await {
let _ = socket.send_to(&buf[..len], from).await;
}
});
addr
}
fn wrap(dest: SocketAddr, payload: &[u8]) -> Vec<u8> {
let mut out = vec![0x00, 0x00, 0x00, 0x01];
match dest.ip() {
IpAddr::V4(ip) => out.extend_from_slice(&ip.octets()),
IpAddr::V6(_) => unreachable!("test uses ipv4"),
}
out.extend_from_slice(&dest.port().to_be_bytes());
out.extend_from_slice(payload);
out
}
#[tokio::test]
async fn relayed_datagrams_are_delivered_and_counted_both_ways() {
let destination = udp_echo().await;
let registry = Arc::new(Registry::new());
let relay = Relay::bind(Arc::clone(®istry), None).await.unwrap();
let relay_addr = relay.local_addr().unwrap();
tokio::spawn(relay.run());
let client = UdpSocket::bind(("127.0.0.1", 0)).await.unwrap();
let payload = b"query bytes";
client
.send_to(&wrap(destination, payload), relay_addr)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
let (len, _) = client.recv_from(&mut buf).await.unwrap();
let echoed = socks::parse_datagram(&buf[..len]).unwrap();
assert_eq!(echoed.payload, payload, "relay must be byte-transparent");
assert_eq!(echoed.port, destination.port(), "reply names its source");
let stats = registry.endpoint("127.0.0.1");
assert_eq!(stats.egress(), payload.len() as u64);
assert_eq!(stats.ingress(), payload.len() as u64);
assert_eq!(stats.connections(), 1);
}
#[tokio::test]
async fn many_datagrams_accumulate_under_one_endpoint() {
let destination = udp_echo().await;
let registry = Arc::new(Registry::new());
let relay = Relay::bind(Arc::clone(®istry), None).await.unwrap();
let relay_addr = relay.local_addr().unwrap();
tokio::spawn(relay.run());
let client = UdpSocket::bind(("127.0.0.1", 0)).await.unwrap();
let payload = [0xEEu8; 512];
let mut buf = vec![0u8; 2048];
for _ in 0..5 {
client
.send_to(&wrap(destination, &payload), relay_addr)
.await
.unwrap();
client.recv_from(&mut buf).await.unwrap();
}
let stats = registry.endpoint("127.0.0.1");
assert_eq!(stats.egress(), 5 * 512);
assert_eq!(stats.ingress(), 5 * 512);
assert_eq!(
stats.connections(),
1,
"one destination is one row, however many datagrams"
);
}
#[tokio::test]
async fn malformed_datagrams_are_dropped_without_killing_the_relay() {
let destination = udp_echo().await;
let registry = Arc::new(Registry::new());
let relay = Relay::bind(Arc::clone(®istry), None).await.unwrap();
let relay_addr = relay.local_addr().unwrap();
tokio::spawn(relay.run());
let client = UdpSocket::bind(("127.0.0.1", 0)).await.unwrap();
client
.send_to(&[0x00, 0x00, 0x07, 0x01, 1, 2, 3, 4, 0, 53], relay_addr)
.await
.unwrap();
client.send_to(&[0x00, 0x00], relay_addr).await.unwrap();
client
.send_to(&wrap(destination, b"ok"), relay_addr)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
let (len, _) = client.recv_from(&mut buf).await.unwrap();
assert_eq!(socks::parse_datagram(&buf[..len]).unwrap().payload, b"ok");
assert_eq!(registry.endpoint("127.0.0.1").egress(), 2);
}
#[tokio::test]
async fn an_unresolvable_destination_is_skipped_without_killing_the_relay() {
let destination = udp_echo().await;
let registry = Arc::new(Registry::new());
let relay = Relay::bind(Arc::clone(®istry), None).await.unwrap();
let relay_addr = relay.local_addr().unwrap();
tokio::spawn(relay.run());
let client = UdpSocket::bind(("127.0.0.1", 0)).await.unwrap();
let unresolvable = [0x00, 0x00, 0x00, 0x03, 0x00, 0x00, 0x35, b'q'];
client.send_to(&unresolvable, relay_addr).await.unwrap();
client
.send_to(&wrap(destination, b"ok"), relay_addr)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
let (len, _) = client.recv_from(&mut buf).await.unwrap();
assert_eq!(
socks::parse_datagram(&buf[..len]).unwrap().payload,
b"ok",
"the relay survives to deliver the next datagram"
);
assert_eq!(
registry.snapshot(None).len(),
1,
"an unresolvable host never becomes a row, since no bytes ever left"
);
assert_eq!(registry.endpoint("127.0.0.1").egress(), 2);
}
}