use std::collections::HashMap;
use std::future::Future;
use std::io;
use std::net::{IpAddr, SocketAddr};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::net::{TcpStream, UdpSocket};
use tracing::debug;
use crate::client_stream::ClientStream;
use crate::config::DnsPolicy;
use crate::dns::{self, DnsResolver};
use crate::errors::Result;
use crate::metrics::Metrics;
use crate::socks5::{self, TargetAddr};
use crate::throttle::Throttle;
const TCP_BUF: usize = 32 * 1024;
const MAX_SHAPED_NAP: Duration = Duration::from_secs(3600);
const UDP_BUF: usize = 65535;
const UDP_SOCKET_BUFFER: usize = 4 * 1024 * 1024;
const MAX_CONSECUTIVE_UDP_RECV_ERRORS: u32 = 16;
const UDP_RECV_BACKOFF_BASE: Duration = Duration::from_millis(1);
const UDP_RECV_BACKOFF_MAX: Duration = Duration::from_millis(100);
pub(crate) fn tune_udp_buffers(socket: &UdpSocket) {
let sock = socket2::SockRef::from(socket);
let _ = sock.set_recv_buffer_size(UDP_SOCKET_BUFFER);
let _ = sock.set_send_buffer_size(UDP_SOCKET_BUFFER);
}
fn is_transient_recv_error(e: &io::Error) -> bool {
matches!(
e.kind(),
io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused
| io::ErrorKind::HostUnreachable
| io::ErrorKind::NetworkUnreachable
| io::ErrorKind::Interrupted
)
}
pub(crate) async fn recv_resilient(
socket: &UdpSocket,
buf: &mut [u8],
consecutive_errors: &mut u32,
) -> io::Result<(usize, SocketAddr)> {
let mut transient_errors: u32 = 0;
loop {
match socket.recv_from(buf).await {
Ok(v) => {
*consecutive_errors = 0;
return Ok(v);
}
Err(e) if is_transient_recv_error(&e) => {
transient_errors = transient_errors.saturating_add(1);
debug!(error = %e, "udp relay recv: ignoring transient error");
tokio::time::sleep(crate::util::capped_exponential_backoff(
transient_errors,
UDP_RECV_BACKOFF_BASE,
UDP_RECV_BACKOFF_MAX,
))
.await;
}
Err(e) => {
*consecutive_errors += 1;
if *consecutive_errors >= MAX_CONSECUTIVE_UDP_RECV_ERRORS {
return Err(e);
}
debug!(error = %e, count = *consecutive_errors, "udp relay recv error; continuing");
tokio::time::sleep(crate::util::capped_exponential_backoff(
*consecutive_errors,
UDP_RECV_BACKOFF_BASE,
UDP_RECV_BACKOFF_MAX,
))
.await;
}
}
}
}
pub(crate) fn client_source_accepted(
src: SocketAddr,
client_ip: IpAddr,
locked: Option<SocketAddr>,
) -> bool {
let client_ip = client_ip.to_canonical();
let src_canon_ip = src.ip().to_canonical();
if src_canon_ip != client_ip {
return false; }
if let Some(locked) = locked {
if locked.ip().to_canonical() != src_canon_ip || locked.port() != src.port() {
return false; }
}
true
}
pub(crate) fn parse_client_header(buf: &[u8]) -> Option<socks5::UdpHeader> {
if buf.get(2).is_some_and(|&frag| frag != 0) {
return None; }
let header = socks5::parse_udp_header(buf).ok()?;
debug_assert_eq!(header.frag, 0, "FRAG must be zero after the raw pre-check");
Some(header)
}
pub struct UdpAssociateOptions {
pub client_ip: IpAddr,
pub client_endpoint: Option<SocketAddr>,
pub idle: Duration,
pub dns_policy: DnsPolicy,
pub dns_resolver: Arc<DnsResolver>,
pub metrics: Arc<Metrics>,
pub throttle: Option<Throttle>,
pub outbound_dual: bool,
pub strict_reply: bool,
}
pub async fn relay_tcp(
client: ClientStream,
remote: TcpStream,
idle: Duration,
throttle: Option<Throttle>,
) -> io::Result<(u64, u64)> {
let (rr, rw) = remote.into_split();
let idle_opt = if idle.is_zero() { None } else { Some(idle) };
match client {
ClientStream::Tcp(client) => {
let (cr, cw) = client.into_split();
relay_streams(cr, cw, rr, rw, idle_opt, throttle).await
}
ClientStream::Tls(client) => {
let (cr, cw) = tokio::io::split(client);
relay_streams(cr, cw, rr, rw, idle_opt, throttle).await
}
}
}
#[cfg(feature = "plugins")]
pub async fn relay_generic<C, R>(
client: C,
remote: R,
idle: Duration,
throttle: Option<Throttle>,
) -> io::Result<(u64, u64)>
where
C: AsyncRead + AsyncWrite + Unpin,
R: AsyncRead + AsyncWrite + Unpin,
{
let (cr, cw) = tokio::io::split(client);
let (rr, rw) = tokio::io::split(remote);
let idle_opt = if idle.is_zero() { None } else { Some(idle) };
relay_streams(cr, cw, rr, rw, idle_opt, throttle).await
}
async fn relay_streams<CR, CW, RR, RW>(
mut client_read: CR,
mut client_write: CW,
mut remote_read: RR,
mut remote_write: RW,
idle: Option<Duration>,
throttle: Option<Throttle>,
) -> io::Result<(u64, u64)>
where
CR: AsyncRead + Unpin,
CW: AsyncWrite + Unpin,
RR: AsyncRead + Unpin,
RW: AsyncWrite + Unpin,
{
let up_total = AtomicU64::new(0);
let down_total = AtomicU64::new(0);
let activity = ActivityClock::new();
let result = {
let up = copy_direction(
&mut client_read,
&mut remote_write,
throttle.as_ref(),
&activity,
&up_total,
idle,
);
let down = copy_direction(
&mut remote_read,
&mut client_write,
throttle.as_ref(),
&activity,
&down_total,
idle,
);
relay_both(up, down, idle, &activity).await
};
let _ = remote_write.shutdown().await;
let _ = client_write.shutdown().await;
result?;
Ok((
up_total.load(Ordering::Relaxed),
down_total.load(Ordering::Relaxed),
))
}
async fn relay_both<U, D>(
up: U,
down: D,
idle: Option<Duration>,
activity: &ActivityClock,
) -> io::Result<()>
where
U: Future<Output = io::Result<()>>,
D: Future<Output = io::Result<()>>,
{
tokio::pin!(up, down);
let mut up_done = false;
let mut down_done = false;
let mut ticker = idle.map(|idle| {
let mut tick = tokio::time::interval(idle_tick_period(idle));
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
tick
});
loop {
tokio::select! {
biased;
res = &mut up, if !up_done => {
res?; up_done = true;
}
res = &mut down, if !down_done => {
res?; down_done = true;
}
_ = maybe_tick(ticker.as_mut()) => {
if idle.is_some_and(|idle| activity.idle_for() >= idle) {
break; }
continue;
}
}
if up_done && down_done {
break;
}
}
Ok(())
}
async fn maybe_tick(ticker: Option<&mut tokio::time::Interval>) {
match ticker {
Some(tick) => {
tick.tick().await;
}
None => std::future::pending::<()>().await,
}
}
async fn copy_direction<R, W>(
r: &mut R,
w: &mut W,
throttle: Option<&Throttle>,
activity: &ActivityClock,
total: &AtomicU64,
idle: Option<Duration>,
) -> io::Result<()>
where
R: AsyncRead + Unpin,
W: AsyncWrite + Unpin,
{
let mut buf = vec![0u8; TCP_BUF];
loop {
let n = r.read(&mut buf).await?;
if n == 0 {
break;
}
activity.mark();
if let Some(throttle) = throttle {
shaped_wait(throttle.reserve(n as u64), activity, idle).await;
}
w.write_all(&buf[..n]).await?;
total.fetch_add(n as u64, Ordering::Relaxed);
}
let _ = w.shutdown().await;
Ok(())
}
async fn shaped_wait(wait: Duration, activity: &ActivityClock, idle: Option<Duration>) {
if wait.is_zero() {
return;
}
let cap = idle
.map(|d| d / 2)
.filter(|s| !s.is_zero())
.map_or(MAX_SHAPED_NAP, |s| s.min(MAX_SHAPED_NAP));
let mut remaining = wait;
loop {
let nap = remaining.min(cap);
tokio::time::sleep(nap).await;
activity.mark();
remaining = remaining.saturating_sub(nap);
if remaining.is_zero() {
break;
}
}
}
pub async fn run_udp_associate<C, F, G>(
mut control: C,
relay_socket: UdpSocket,
outbound: UdpSocket,
options: UdpAssociateOptions,
authorize: F,
on_datagram: G,
) -> Result<()>
where
C: AsyncRead + AsyncWrite + Unpin,
F: Fn(Option<&str>, IpAddr, u16) -> bool + Send + Sync + 'static,
G: Fn(bool, SocketAddr, &[u8]) -> bool + Clone + Send + Sync + 'static,
{
let relay_socket = Arc::new(relay_socket);
let outbound = Arc::new(outbound);
let client_endpoint = Arc::new(OnceLock::new());
if let Some(endpoint) = options.client_endpoint {
let _ = client_endpoint.set(endpoint);
}
let activity = Arc::new(ActivityClock::new());
let contacted = Arc::new(Mutex::new(ContactedRemotes::new(options.strict_reply)));
let mut client_to_remote = tokio::spawn(relay_client_to_remote(
relay_socket.clone(),
outbound.clone(),
options.client_ip,
options.dns_policy,
options.dns_resolver,
options.metrics.clone(),
options.throttle.clone(),
client_endpoint.clone(),
activity.clone(),
contacted.clone(),
options.outbound_dual,
authorize,
on_datagram.clone(),
));
let mut remote_to_client = tokio::spawn(relay_remote_to_client(
outbound,
relay_socket,
options.metrics,
options.throttle,
client_endpoint,
activity.clone(),
contacted,
on_datagram,
));
let idle_enabled = !options.idle.is_zero();
let tick_period = if idle_enabled {
idle_tick_period(options.idle)
} else {
Duration::from_secs(3600)
};
let mut idle_tick = tokio::time::interval(tick_period);
idle_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let mut ctrl_buf = [0u8; 512];
let result = loop {
tokio::select! {
res = control.read(&mut ctrl_buf) => {
match res {
Ok(0) | Err(_) => break Ok(()),
Ok(_) => continue, }
}
res = &mut client_to_remote => break flatten_direction(res),
res = &mut remote_to_client => break flatten_direction(res),
_ = idle_tick.tick() => {
if idle_enabled && activity.idle_for() >= options.idle {
break Ok(()); }
}
}
};
client_to_remote.abort();
remote_to_client.abort();
result
}
fn flatten_direction(
joined: std::result::Result<Result<()>, tokio::task::JoinError>,
) -> Result<()> {
match joined {
Ok(result) => result,
Err(e) => Err(crate::errors::Error::Io(std::io::Error::other(e))),
}
}
#[cfg(feature = "plugins")]
pub(crate) async fn drive_owned_association<C>(
mut control: C,
relay_socket: UdpSocket,
outbound: UdpSocket,
options: UdpAssociateOptions,
authorize: DatagramAuthorizer,
interceptor: Box<dyn crate::plugin::DatagramInterceptor>,
) -> Result<()>
where
C: AsyncRead + AsyncWrite + Unpin,
{
let client_endpoint = Arc::new(OnceLock::new());
if let Some(endpoint) = options.client_endpoint {
let _ = client_endpoint.set(endpoint);
}
let activity = Arc::new(ActivityClock::new());
let contacted = Arc::new(Mutex::new(ContactedRemotes::new(options.strict_reply)));
let client = ClientDatagrams::new(
Arc::new(relay_socket),
options.client_ip,
client_endpoint,
activity.clone(),
options.throttle.clone(),
options.metrics.clone(),
);
let upstream = UpstreamOriginator::new(
Arc::new(outbound),
options.outbound_dual,
options.dns_policy.clone(),
authorize,
contacted,
activity.clone(),
options.throttle.clone(),
options.metrics.clone(),
);
let args = crate::plugin::AssociationArgs::new(client, upstream, options.idle);
let interceptor = interceptor.run(args);
tokio::pin!(interceptor);
let idle_enabled = !options.idle.is_zero();
let tick_period = if idle_enabled {
idle_tick_period(options.idle)
} else {
Duration::from_secs(3600)
};
let mut idle_tick = tokio::time::interval(tick_period);
idle_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let mut ctrl_buf = [0u8; 512];
loop {
tokio::select! {
res = control.read(&mut ctrl_buf) => {
match res {
Ok(0) | Err(_) => break Ok(()),
Ok(_) => continue,
}
}
res = &mut interceptor => break flatten_interceptor(res),
_ = idle_tick.tick() => {
if idle_enabled && activity.idle_for() >= options.idle {
break Ok(());
}
}
}
}
}
#[cfg(feature = "plugins")]
fn flatten_interceptor(res: io::Result<crate::plugin::FlowStats>) -> Result<()> {
match res {
Ok(_stats) => Ok(()),
Err(e) => Err(crate::errors::Error::Io(e)),
}
}
const MAX_CONTACTED_REMOTES: usize = 256;
pub(crate) struct ContactedRemotes {
seen: HashMap<SocketAddr, u64>,
ticks: u64,
match_port: bool,
}
impl ContactedRemotes {
fn new(match_port: bool) -> Self {
ContactedRemotes {
seen: HashMap::new(),
ticks: 0,
match_port,
}
}
fn key(&self, addr: SocketAddr) -> SocketAddr {
if self.match_port {
addr
} else {
SocketAddr::new(addr.ip(), 0)
}
}
fn record(&mut self, addr: SocketAddr) {
let key = self.key(addr);
self.ticks += 1;
if self.seen.len() >= MAX_CONTACTED_REMOTES && !self.seen.contains_key(&key) {
if let Some(oldest) = self
.seen
.iter()
.min_by_key(|(_, &tick)| tick)
.map(|(addr, _)| *addr)
{
self.seen.remove(&oldest);
}
}
self.seen.insert(key, self.ticks);
}
fn contains(&self, addr: SocketAddr) -> bool {
self.seen.contains_key(&self.key(addr))
}
}
pub(crate) struct ActivityClock {
started: tokio::time::Instant,
last_ms: AtomicU64,
}
impl ActivityClock {
fn new() -> Self {
ActivityClock {
started: tokio::time::Instant::now(),
last_ms: AtomicU64::new(0),
}
}
fn mark(&self) {
let ms = self.started.elapsed().as_millis() as u64;
self.last_ms.store(ms, Ordering::Relaxed);
}
fn idle_for(&self) -> Duration {
let now_ms = self.started.elapsed().as_millis() as u64;
Duration::from_millis(now_ms.saturating_sub(self.last_ms.load(Ordering::Relaxed)))
}
}
fn idle_tick_period(idle: Duration) -> Duration {
(idle / 4).clamp(Duration::from_millis(250), Duration::from_secs(15))
}
#[allow(clippy::too_many_arguments)]
async fn relay_client_to_remote<F, G>(
relay_socket: Arc<UdpSocket>,
outbound: Arc<UdpSocket>,
client_ip: IpAddr,
dns_policy: DnsPolicy,
dns_resolver: Arc<DnsResolver>,
metrics: Arc<Metrics>,
throttle: Option<Throttle>,
client_endpoint: Arc<OnceLock<SocketAddr>>,
activity: Arc<ActivityClock>,
contacted: Arc<Mutex<ContactedRemotes>>,
outbound_dual: bool,
authorize: F,
on_datagram: G,
) -> Result<()>
where
F: Fn(Option<&str>, IpAddr, u16) -> bool,
G: Fn(bool, SocketAddr, &[u8]) -> bool,
{
let mut buf = vec![0u8; UDP_BUF];
let mut recv_errors = 0u32;
loop {
let (n, src) = recv_resilient(&relay_socket, &mut buf, &mut recv_errors).await?;
if !client_source_accepted(src, client_ip, client_endpoint.get().copied()) {
continue; }
let header = match parse_client_header(&buf[..n]) {
Some(h) => h,
None => continue, };
let dest = match &header.dest {
TargetAddr::Ip(sa) => {
let sa = SocketAddr::new(sa.ip().to_canonical(), sa.port());
if !dns::address_allowed(sa.ip(), &dns_policy) {
continue;
}
sa
}
domain => match dns_resolver.resolve_one(domain, &dns_policy).await {
Ok(Some(sa)) => sa,
Ok(None) => continue,
Err(_) => continue,
},
};
let host = match &header.dest {
TargetAddr::Domain(d, _) => Some(d.as_str()),
TargetAddr::Ip(_) => None,
};
if !authorize(host, dest.ip(), dest.port()) {
continue;
}
if !on_datagram(false, dest, &buf[header.payload_offset..n]) {
continue;
}
let _ = client_endpoint.set(src);
let dest_canon = SocketAddr::new(dest.ip().to_canonical(), dest.port());
contacted
.lock()
.unwrap_or_else(|e| e.into_inner())
.record(dest_canon);
activity.mark();
let payload = &buf[header.payload_offset..n];
if throttle
.as_ref()
.is_some_and(|t| !t.police(payload.len() as u64))
{
metrics.rate_limited();
continue;
}
let mut send_dest = dest;
if outbound_dual {
if let SocketAddr::V4(v4) = dest {
send_dest.set_ip(IpAddr::V6(v4.ip().to_ipv6_mapped()));
}
}
match outbound.send_to(payload, send_dest).await {
Ok(_) => metrics.udp_client_packet_relayed(payload.len() as u64),
Err(e) => {
metrics.udp_send_failed();
debug!(dest = %dest, error = %e, "UDP outbound send failed");
}
}
}
}
#[allow(clippy::too_many_arguments)]
async fn relay_remote_to_client<G>(
outbound: Arc<UdpSocket>,
relay_socket: Arc<UdpSocket>,
metrics: Arc<Metrics>,
throttle: Option<Throttle>,
client_endpoint: Arc<OnceLock<SocketAddr>>,
activity: Arc<ActivityClock>,
contacted: Arc<Mutex<ContactedRemotes>>,
on_datagram: G,
) -> Result<()>
where
G: Fn(bool, SocketAddr, &[u8]) -> bool,
{
let mut buf = vec![0u8; socks5::UDP_IP_HEADER_MAX + UDP_BUF];
let mut recv_errors = 0u32;
loop {
let (n, remote_src) = recv_resilient(
&outbound,
&mut buf[socks5::UDP_IP_HEADER_MAX..],
&mut recv_errors,
)
.await?;
let Some(caddr) = client_endpoint.get().copied() else {
continue; };
let remote_canon = SocketAddr::new(remote_src.ip().to_canonical(), remote_src.port());
if !contacted
.lock()
.unwrap_or_else(|e| e.into_inner())
.contains(remote_canon)
{
continue;
}
if !on_datagram(
true,
remote_canon,
&buf[socks5::UDP_IP_HEADER_MAX..socks5::UDP_IP_HEADER_MAX + n],
) {
continue;
}
activity.mark();
let prefix: &mut [u8; socks5::UDP_IP_HEADER_MAX] = (&mut buf[..socks5::UDP_IP_HEADER_MAX])
.try_into()
.expect("prefix slice is UDP_IP_HEADER_MAX bytes");
let start = socks5::write_udp_header_tail(remote_src, prefix);
let datagram = &buf[start..socks5::UDP_IP_HEADER_MAX + n];
if throttle
.as_ref()
.is_some_and(|t| !t.police(datagram.len() as u64))
{
metrics.rate_limited();
continue;
}
match relay_socket.send_to(datagram, caddr).await {
Ok(_) => metrics.udp_remote_packet_relayed(n as u64),
Err(e) => {
metrics.udp_send_failed();
debug!(client = %caddr, error = %e, "UDP reply to client failed");
}
}
}
}
#[cfg(feature = "plugins")]
pub struct ClientDatagrams {
socket: Arc<UdpSocket>,
client_ip: IpAddr,
client_endpoint: Arc<OnceLock<SocketAddr>>,
activity: Arc<ActivityClock>,
throttle: Option<Throttle>,
metrics: Arc<Metrics>,
}
#[cfg(feature = "plugins")]
impl ClientDatagrams {
pub(crate) fn new(
socket: Arc<UdpSocket>,
client_ip: IpAddr,
client_endpoint: Arc<OnceLock<SocketAddr>>,
activity: Arc<ActivityClock>,
throttle: Option<Throttle>,
metrics: Arc<Metrics>,
) -> Self {
ClientDatagrams {
socket,
client_ip,
client_endpoint,
activity,
throttle,
metrics,
}
}
pub fn for_interceptor(
socket: Arc<UdpSocket>,
client_ip: IpAddr,
client_endpoint: Option<SocketAddr>,
) -> Self {
let lock = Arc::new(OnceLock::new());
if let Some(endpoint) = client_endpoint {
let _ = lock.set(endpoint);
}
ClientDatagrams::new(
socket,
client_ip,
lock,
Arc::new(ActivityClock::new()),
None,
Metrics::new(),
)
}
pub async fn recv(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
let mut recv_errors = 0u32;
loop {
let (n, src) = recv_resilient(&self.socket, buf, &mut recv_errors).await?;
if !client_source_accepted(src, self.client_ip, self.client_endpoint.get().copied()) {
continue; }
if buf[..n].get(3) == Some(&socks5::ATYP_DOMAIN) {
continue;
}
let header = match parse_client_header(&buf[..n]) {
Some(h) => h,
None => continue, };
let payload_offset = header.payload_offset;
let origin = match header.dest {
TargetAddr::Ip(sa) => SocketAddr::new(sa.ip().to_canonical(), sa.port()),
TargetAddr::Domain(..) => continue,
};
let _ = self.client_endpoint.set(src);
self.activity.mark();
let payload_len = n - payload_offset;
buf.copy_within(payload_offset..n, 0);
return Ok((payload_len, origin));
}
}
pub async fn send(&self, from: SocketAddr, payload: &[u8]) -> io::Result<()> {
let Some(caddr) = self.client_endpoint.get().copied() else {
return Ok(()); };
let mut datagram = vec![0u8; socks5::UDP_IP_HEADER_MAX + payload.len()];
datagram[socks5::UDP_IP_HEADER_MAX..].copy_from_slice(payload);
let prefix: &mut [u8; socks5::UDP_IP_HEADER_MAX] = (&mut datagram
[..socks5::UDP_IP_HEADER_MAX])
.try_into()
.expect("prefix slice is UDP_IP_HEADER_MAX bytes");
let start = socks5::write_udp_header_tail(from, prefix);
let out = &datagram[start..];
if self
.throttle
.as_ref()
.is_some_and(|t| !t.police(out.len() as u64))
{
self.metrics.rate_limited();
return Ok(());
}
self.socket.send_to(out, caddr).await?;
self.metrics.udp_remote_packet_relayed(payload.len() as u64);
Ok(())
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.socket.local_addr()
}
pub fn client_endpoint(&self) -> Option<SocketAddr> {
self.client_endpoint.get().copied()
}
}
#[cfg(feature = "plugins")]
#[non_exhaustive]
pub struct UpstreamTarget {
dst: SocketAddr,
}
#[cfg(feature = "plugins")]
impl UpstreamTarget {
pub fn dst(&self) -> SocketAddr {
self.dst
}
}
#[cfg(feature = "plugins")]
pub type DatagramAuthorizer = Arc<dyn Fn(Option<&str>, IpAddr, u16) -> bool + Send + Sync>;
#[cfg(feature = "plugins")]
pub struct UpstreamOriginator {
outbound: Arc<UdpSocket>,
outbound_dual: bool,
dns_policy: DnsPolicy,
authorize: DatagramAuthorizer,
contacted: Arc<Mutex<ContactedRemotes>>,
activity: Arc<ActivityClock>,
throttle: Option<Throttle>,
metrics: Arc<Metrics>,
}
#[cfg(feature = "plugins")]
impl UpstreamOriginator {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
outbound: Arc<UdpSocket>,
outbound_dual: bool,
dns_policy: DnsPolicy,
authorize: DatagramAuthorizer,
contacted: Arc<Mutex<ContactedRemotes>>,
activity: Arc<ActivityClock>,
throttle: Option<Throttle>,
metrics: Arc<Metrics>,
) -> Self {
UpstreamOriginator {
outbound,
outbound_dual,
dns_policy,
authorize,
contacted,
activity,
throttle,
metrics,
}
}
pub fn for_interceptor(
outbound: Arc<UdpSocket>,
outbound_dual: bool,
strict_reply: bool,
authorize: DatagramAuthorizer,
) -> Self {
UpstreamOriginator::new(
outbound,
outbound_dual,
interceptor_dns_policy(),
authorize,
Arc::new(Mutex::new(ContactedRemotes::new(strict_reply))),
Arc::new(ActivityClock::new()),
None,
Metrics::new(),
)
}
pub fn authorize(&self, host: Option<&str>, dst: SocketAddr) -> Option<UpstreamTarget> {
let dst = SocketAddr::new(dst.ip().to_canonical(), dst.port());
if !dns::address_allowed(dst.ip(), &self.dns_policy) {
return None;
}
if !(self.authorize)(host, dst.ip(), dst.port()) {
return None;
}
Some(UpstreamTarget { dst })
}
pub async fn send_to(&self, target: &UpstreamTarget, payload: &[u8]) -> io::Result<()> {
self.contacted
.lock()
.unwrap_or_else(|e| e.into_inner())
.record(target.dst);
if self
.throttle
.as_ref()
.is_some_and(|t| !t.police(payload.len() as u64))
{
self.metrics.rate_limited();
return Ok(());
}
let mut send_dest = target.dst;
if self.outbound_dual {
if let SocketAddr::V4(v4) = target.dst {
send_dest.set_ip(IpAddr::V6(v4.ip().to_ipv6_mapped()));
}
}
self.outbound.send_to(payload, send_dest).await?;
self.metrics.udp_client_packet_relayed(payload.len() as u64);
Ok(())
}
pub async fn recv(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
let mut recv_errors = 0u32;
loop {
let (n, remote_src) = recv_resilient(&self.outbound, buf, &mut recv_errors).await?;
let remote_canon = SocketAddr::new(remote_src.ip().to_canonical(), remote_src.port());
if !self
.contacted
.lock()
.unwrap_or_else(|e| e.into_inner())
.contains(remote_canon)
{
continue; }
self.activity.mark();
return Ok((n, remote_canon));
}
}
pub fn is_dual_stack(&self) -> bool {
self.outbound_dual
}
}
#[cfg(feature = "plugins")]
fn interceptor_dns_policy() -> DnsPolicy {
DnsPolicy {
preference: crate::config::DnsPreference::System,
try_all: false,
deny: Vec::new(),
cache_ttl: None,
timeout: Duration::from_secs(5),
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
#[tokio::test(start_paused = true)]
async fn idle_relay_stays_alive_while_one_direction_flows() {
let (client, mut client_peer) = tokio::io::duplex(64 * 1024);
let (remote, mut remote_peer) = tokio::io::duplex(64 * 1024);
let (cr, cw) = tokio::io::split(client);
let (rr, rw) = tokio::io::split(remote);
let relay = tokio::spawn(relay_streams(
cr,
cw,
rr,
rw,
Some(Duration::from_secs(5)),
None,
));
let mut received = [0u8; 4];
for _ in 0..3 {
tokio::time::sleep(Duration::from_secs(2)).await;
remote_peer.write_all(b"ping").await.unwrap();
client_peer.read_exact(&mut received).await.unwrap();
assert_eq!(&received, b"ping");
}
assert!(!relay.is_finished());
drop(remote_peer);
drop(client_peer);
let (up, down) = relay.await.unwrap().unwrap();
assert_eq!(up, 0);
assert_eq!(down, 12);
}
#[tokio::test(start_paused = true)]
async fn idle_relay_times_out_quiet_connections() {
let (client, client_peer) = tokio::io::duplex(64 * 1024);
let (remote, remote_peer) = tokio::io::duplex(64 * 1024);
let (cr, cw) = tokio::io::split(client);
let (rr, rw) = tokio::io::split(remote);
let relay = tokio::spawn(relay_streams(
cr,
cw,
rr,
rw,
Some(Duration::from_secs(5)),
None,
));
let result = tokio::time::timeout(Duration::from_secs(30), relay)
.await
.expect("idle watchdog should end the relay");
let (up, down) = result.unwrap().unwrap();
assert_eq!((up, down), (0, 0));
drop(client_peer);
drop(remote_peer);
}
#[tokio::test]
async fn relay_tcp_echoes_both_directions() {
let echo = TcpListener::bind("127.0.0.1:0").await.unwrap();
let echo_addr = echo.local_addr().unwrap();
tokio::spawn(async move {
let (mut s, _) = echo.accept().await.unwrap();
let mut buf = [0u8; 1024];
loop {
let n = s.read(&mut buf).await.unwrap();
if n == 0 {
break;
}
s.write_all(&buf[..n]).await.unwrap();
}
});
let client_side = TcpListener::bind("127.0.0.1:0").await.unwrap();
let client_side_addr = client_side.local_addr().unwrap();
let remote = TcpStream::connect(echo_addr).await.unwrap();
let mut client = TcpStream::connect(client_side_addr).await.unwrap();
let (proxy_client, _) = client_side.accept().await.unwrap();
tokio::spawn(async move {
let _ = relay_tcp(
ClientStream::Tcp(proxy_client),
remote,
Duration::ZERO,
None,
)
.await;
});
client.write_all(b"ping").await.unwrap();
let mut buf = [0u8; 4];
client.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"ping");
}
#[tokio::test(start_paused = true)]
async fn copy_direction_shapes_instead_of_dropping() {
use crate::throttle::TokenBucket;
use std::sync::Mutex;
let throttle = Throttle::new().with_bucket(Arc::new(Mutex::new(TokenBucket::new(
1000.0,
1000.0,
std::time::Instant::now(),
))));
let data = vec![0u8; 3000];
let mut reader = &data[..];
let mut writer = tokio::io::sink();
let activity = ActivityClock::new();
let total = AtomicU64::new(0);
let start = tokio::time::Instant::now();
copy_direction(
&mut reader,
&mut writer,
Some(&throttle),
&activity,
&total,
None,
)
.await
.unwrap();
assert_eq!(total.load(Ordering::Relaxed), 3000, "all bytes forwarded");
assert!(
start.elapsed() >= Duration::from_millis(1900),
"expected ~2s of shaping, got {:?}",
start.elapsed()
);
}
#[tokio::test(start_paused = true)]
async fn relay_both_fails_fast_when_a_direction_errors() {
let activity = ActivityClock::new();
let up = async { Err::<(), io::Error>(io::Error::other("up failed")) };
let down = std::future::pending::<io::Result<()>>();
let outcome = tokio::time::timeout(
Duration::from_secs(3600),
relay_both(up, down, None, &activity),
)
.await
.expect("relay must not hang when one direction errors");
assert!(outcome.is_err(), "the copy error must propagate");
}
#[tokio::test(start_paused = true)]
async fn relay_both_waits_for_the_open_half_after_a_clean_close() {
let activity = ActivityClock::new();
let up = async { Ok::<(), io::Error>(()) };
let down = std::future::pending::<io::Result<()>>();
let outcome = tokio::time::timeout(
Duration::from_secs(60),
relay_both(up, down, None, &activity),
)
.await;
assert!(
outcome.is_err(),
"relay must keep waiting for the open half after a clean half-close"
);
}
#[tokio::test(start_paused = true)]
async fn relay_both_completes_when_both_directions_finish() {
let activity = ActivityClock::new();
let up = async { Ok::<(), io::Error>(()) };
let down = async { Ok::<(), io::Error>(()) };
relay_both(up, down, None, &activity)
.await
.expect("both directions closing cleanly is success");
}
#[tokio::test(start_paused = true)]
async fn relay_both_ends_on_idle_timeout() {
let activity = ActivityClock::new();
let up = std::future::pending::<io::Result<()>>();
let down = std::future::pending::<io::Result<()>>();
let outcome = tokio::time::timeout(
Duration::from_secs(3600),
relay_both(up, down, Some(Duration::from_secs(10)), &activity),
)
.await
.expect("the idle watchdog must end the relay");
assert!(outcome.is_ok(), "an idle timeout ends the relay cleanly");
}
#[tokio::test(start_paused = true)]
async fn shaped_wait_keeps_connection_alive_under_small_idle() {
let activity = ActivityClock::new();
let idle = Duration::from_secs(2);
let wait = shaped_wait(Duration::from_secs(10), &activity, Some(idle));
tokio::pin!(wait);
let mut tick = tokio::time::interval(idle);
tick.tick().await; loop {
tokio::select! {
_ = &mut wait => break,
_ = tick.tick() => {
assert!(
activity.idle_for() < idle,
"connection looked idle for {:?} while shaping",
activity.idle_for()
);
}
}
}
}
fn udp_options(client_ip: IpAddr, idle: Duration) -> UdpAssociateOptions {
UdpAssociateOptions {
client_ip,
client_endpoint: None,
idle,
dns_policy: DnsPolicy {
preference: crate::config::DnsPreference::System,
try_all: false,
deny: Vec::new(),
cache_ttl: None,
timeout: Duration::from_secs(5),
},
dns_resolver: Arc::new(DnsResolver::new()),
metrics: Metrics::new(),
throttle: None,
outbound_dual: false,
strict_reply: false,
}
}
#[tokio::test]
async fn spoofed_source_datagrams_do_not_refresh_udp_idle() {
let (control, _control_peer) = tokio::io::duplex(1024);
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let relay_addr = relay_socket.local_addr().unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let options = udp_options(IpAddr::from([10, 0, 0, 1]), Duration::from_millis(500));
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true,
|_, _, _| true,
));
let sender = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let spray = tokio::spawn(async move {
for _ in 0..40 {
let _ = sender.send_to(b"junk", relay_addr).await;
tokio::time::sleep(Duration::from_millis(100)).await;
}
});
let result = tokio::time::timeout(Duration::from_secs(2), assoc)
.await
.expect("association should idle out despite spoofed traffic");
assert!(result.unwrap().is_ok());
spray.abort();
}
#[tokio::test]
async fn control_channel_data_does_not_refresh_udp_idle() {
let (control, mut control_peer) = tokio::io::duplex(1024);
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let options = udp_options(IpAddr::from([127, 0, 0, 1]), Duration::from_millis(500));
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true,
|_, _, _| true,
));
let junk = tokio::spawn(async move {
for _ in 0..40 {
if control_peer.write_all(b"x").await.is_err() {
break;
}
tokio::time::sleep(Duration::from_millis(100)).await;
}
});
let result = tokio::time::timeout(Duration::from_secs(2), assoc)
.await
.expect("association should idle out despite control-channel data");
assert!(result.unwrap().is_ok());
junk.abort();
}
#[tokio::test]
async fn validated_datagrams_keep_udp_association_alive() {
let (control, _control_peer) = tokio::io::duplex(1024);
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let relay_addr = relay_socket.local_addr().unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let options = udp_options(IpAddr::from([127, 0, 0, 1]), Duration::from_millis(500));
let mut assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true,
|_, _, _| true,
));
let sender = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let mut datagram = vec![0u8, 0, 0, 1];
datagram.extend_from_slice(&[127, 0, 0, 1]);
datagram.extend_from_slice(&9u16.to_be_bytes());
datagram.extend_from_slice(b"ping");
let keepalive = tokio::spawn(async move {
for _ in 0..20 {
let _ = sender.send_to(&datagram, relay_addr).await;
tokio::time::sleep(Duration::from_millis(100)).await;
}
});
let still_running = tokio::time::timeout(Duration::from_millis(1500), &mut assoc).await;
assert!(
still_running.is_err(),
"validated traffic must keep the association alive"
);
keepalive.abort();
assoc.abort();
}
#[test]
fn contacted_remotes_records_and_evicts_oldest() {
let mut contacted = ContactedRemotes::new(false);
let addr = |i: usize| SocketAddr::new(IpAddr::from([10, 0, (i >> 8) as u8, i as u8]), 9);
contacted.record(addr(0));
assert!(contacted.contains(addr(0)));
assert!(!contacted.contains(addr(1)));
for i in 1..MAX_CONTACTED_REMOTES {
contacted.record(addr(i));
}
assert!(contacted.contains(addr(0)));
contacted.record(addr(MAX_CONTACTED_REMOTES));
assert!(
!contacted.contains(addr(0)),
"the oldest contacted remote should be evicted at the cap"
);
assert!(contacted.contains(addr(MAX_CONTACTED_REMOTES)));
assert!(contacted.contains(addr(MAX_CONTACTED_REMOTES - 1)));
}
#[test]
fn contacted_remotes_port_matching_modes() {
let dest = SocketAddr::from(([10, 0, 0, 1], 53));
let same_ip_other_port = SocketAddr::from(([10, 0, 0, 1], 9999));
let other_ip = SocketAddr::from(([10, 0, 0, 2], 53));
let mut loose = ContactedRemotes::new(false);
loose.record(dest);
assert!(loose.contains(dest));
assert!(loose.contains(same_ip_other_port));
assert!(!loose.contains(other_ip));
let mut strict = ContactedRemotes::new(true);
strict.record(dest);
assert!(strict.contains(dest));
assert!(!strict.contains(same_ip_other_port));
assert!(!strict.contains(other_ip));
}
#[tokio::test]
async fn udp_drops_replies_from_uncontacted_remotes() {
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let outbound_addr = outbound.local_addr().unwrap();
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client.local_addr().unwrap();
let (control, _control_peer) = tokio::io::duplex(1024);
let mut options = udp_options(client_addr.ip(), Duration::from_secs(5));
options.client_endpoint = Some(client_addr);
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true,
|_, _, _| true,
));
let stranger = UdpSocket::bind("127.0.0.1:0").await.unwrap();
stranger
.send_to(b"unsolicited", outbound_addr)
.await
.unwrap();
let mut buf = [0u8; 256];
assert!(
tokio::time::timeout(Duration::from_millis(500), client.recv_from(&mut buf))
.await
.is_err(),
"an unsolicited reply must not reach the client"
);
assoc.abort();
}
#[tokio::test]
async fn udp_forwards_replies_from_contacted_remotes() {
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let relay_addr = relay_socket.local_addr().unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let outbound_addr = outbound.local_addr().unwrap();
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client.local_addr().unwrap();
let dest = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let dest_addr = dest.local_addr().unwrap();
let (control, _control_peer) = tokio::io::duplex(1024);
let mut options = udp_options(client_addr.ip(), Duration::from_secs(5));
options.client_endpoint = Some(client_addr);
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true,
|_, _, _| true,
));
let IpAddr::V4(dest_ip) = dest_addr.ip() else {
unreachable!()
};
let mut datagram = vec![0u8, 0, 0, 1]; datagram.extend_from_slice(&dest_ip.octets());
datagram.extend_from_slice(&dest_addr.port().to_be_bytes());
datagram.extend_from_slice(b"ping");
client.send_to(&datagram, relay_addr).await.unwrap();
let mut dbuf = [0u8; 64];
let (dn, _) = tokio::time::timeout(Duration::from_secs(1), dest.recv_from(&mut dbuf))
.await
.expect("dest should receive the forwarded datagram")
.unwrap();
assert_eq!(&dbuf[..dn], b"ping");
dest.send_to(b"pong", outbound_addr).await.unwrap();
let mut buf = [0u8; 256];
let (cn, _) = tokio::time::timeout(Duration::from_secs(1), client.recv_from(&mut buf))
.await
.expect("a reply from a contacted remote must reach the client")
.unwrap();
assert!(buf[..cn].ends_with(b"pong"));
assoc.abort();
}
#[tokio::test]
async fn udp_datagram_to_denied_destination_is_dropped() {
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let relay_addr = relay_socket.local_addr().unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client.local_addr().unwrap();
let dest = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let dest_addr = dest.local_addr().unwrap();
let (control, _control_peer) = tokio::io::duplex(1024);
let mut options = udp_options(client_addr.ip(), Duration::from_secs(5));
options.client_endpoint = Some(client_addr);
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| false,
|_, _, _| true,
));
let IpAddr::V4(dest_ip) = dest_addr.ip() else {
unreachable!()
};
let mut datagram = vec![0u8, 0, 0, 1]; datagram.extend_from_slice(&dest_ip.octets());
datagram.extend_from_slice(&dest_addr.port().to_be_bytes());
datagram.extend_from_slice(b"ping");
client.send_to(&datagram, relay_addr).await.unwrap();
let mut dbuf = [0u8; 64];
let forwarded =
tokio::time::timeout(Duration::from_secs(1), dest.recv_from(&mut dbuf)).await;
assert!(
forwarded.is_err(),
"a datagram the authorizer denied must not be forwarded to the destination"
);
assert!(
!assoc.is_finished(),
"association task exited prematurely; the no-forward result is not meaningful"
);
assoc.abort();
}
#[tokio::test]
async fn on_datagram_can_drop_client_to_target() {
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let relay_addr = relay_socket.local_addr().unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client.local_addr().unwrap();
let dest = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let dest_addr = dest.local_addr().unwrap();
let (control, _control_peer) = tokio::io::duplex(1024);
let mut options = udp_options(client_addr.ip(), Duration::from_secs(5));
options.client_endpoint = Some(client_addr);
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true, |is_reply, _, _| is_reply, ));
let IpAddr::V4(dest_ip) = dest_addr.ip() else {
unreachable!()
};
let mut datagram = vec![0u8, 0, 0, 1];
datagram.extend_from_slice(&dest_ip.octets());
datagram.extend_from_slice(&dest_addr.port().to_be_bytes());
datagram.extend_from_slice(b"ping");
client.send_to(&datagram, relay_addr).await.unwrap();
let mut dbuf = [0u8; 64];
let forwarded =
tokio::time::timeout(Duration::from_secs(1), dest.recv_from(&mut dbuf)).await;
assert!(
forwarded.is_err(),
"a datagram dropped by on_datagram must not be forwarded to the destination"
);
assert!(
!assoc.is_finished(),
"association task exited prematurely; the no-forward result is not meaningful"
);
assoc.abort();
}
#[tokio::test]
async fn on_datagram_can_drop_target_to_client() {
let relay_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let relay_addr = relay_socket.local_addr().unwrap();
let outbound = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let outbound_addr = outbound.local_addr().unwrap();
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client.local_addr().unwrap();
let dest = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let dest_addr = dest.local_addr().unwrap();
let (control, _control_peer) = tokio::io::duplex(1024);
let mut options = udp_options(client_addr.ip(), Duration::from_secs(5));
options.client_endpoint = Some(client_addr);
let assoc = tokio::spawn(run_udp_associate(
control,
relay_socket,
outbound,
options,
|_, _, _| true, |is_reply, _, _| !is_reply, ));
let IpAddr::V4(dest_ip) = dest_addr.ip() else {
unreachable!()
};
let mut datagram = vec![0u8, 0, 0, 1];
datagram.extend_from_slice(&dest_ip.octets());
datagram.extend_from_slice(&dest_addr.port().to_be_bytes());
datagram.extend_from_slice(b"ping");
client.send_to(&datagram, relay_addr).await.unwrap();
let mut dbuf = [0u8; 64];
let (dn, _) = tokio::time::timeout(Duration::from_secs(1), dest.recv_from(&mut dbuf))
.await
.expect("dest should receive the forwarded datagram")
.unwrap();
assert_eq!(&dbuf[..dn], b"ping");
dest.send_to(b"pong", outbound_addr).await.unwrap();
let mut buf = [0u8; 256];
assert!(
tokio::time::timeout(Duration::from_millis(500), client.recv_from(&mut buf))
.await
.is_err(),
"a reply dropped by on_datagram must not reach the client"
);
assoc.abort();
}
fn sa(s: &str) -> SocketAddr {
s.parse().unwrap()
}
#[test]
fn client_source_rejects_a_different_ip() {
let client_ip = sa("203.0.113.7:0").ip().to_canonical();
assert!(client_source_accepted(
sa("203.0.113.7:5000"),
client_ip,
None
));
assert!(!client_source_accepted(
sa("203.0.113.8:5000"),
client_ip,
None
));
}
#[test]
fn client_source_matches_across_v4_mapped_v6() {
let client_ip = sa("203.0.113.7:0").ip().to_canonical();
assert!(client_source_accepted(
sa("[::ffff:203.0.113.7]:5000"),
client_ip,
None
));
let mapped_client = sa("[::ffff:203.0.113.7]:0").ip(); assert!(client_source_accepted(
sa("203.0.113.7:5000"),
mapped_client,
None
));
}
#[test]
fn client_source_enforces_the_endpoint_lock() {
let client_ip = sa("203.0.113.7:0").ip().to_canonical();
let lock = sa("203.0.113.7:5000");
assert!(client_source_accepted(
sa("203.0.113.7:5000"),
client_ip,
Some(lock)
));
assert!(!client_source_accepted(
sa("203.0.113.7:6000"),
client_ip,
Some(lock)
));
assert!(client_source_accepted(
sa("[::ffff:203.0.113.7]:5000"),
client_ip,
Some(lock)
));
}
#[test]
fn parse_client_header_drops_fragments_and_garbage() {
let mut good = vec![0x00, 0x00, 0x00, 0x01, 203, 0, 113, 9, 0x01, 0xbb];
good.extend_from_slice(b"payload");
let header = parse_client_header(&good).expect("valid header parses");
assert_eq!(header.frag, 0);
assert_eq!(&good[header.payload_offset..], b"payload");
let mut fragmented = good.clone();
fragmented[2] = 0x01;
assert!(parse_client_header(&fragmented).is_none());
assert!(parse_client_header(&[0xff, 0xff, 0x00, 0x01]).is_none());
}
}
#[cfg(all(test, feature = "plugins"))]
mod facade_tests {
use super::*;
use crate::config::{DnsDenyCategory, DnsPreference};
fn permissive_policy() -> DnsPolicy {
DnsPolicy {
preference: DnsPreference::System,
try_all: false,
deny: vec![],
cache_ttl: None,
timeout: Duration::from_secs(5),
}
}
fn socks_dg(dest: SocketAddr, payload: &[u8]) -> Vec<u8> {
let mut dg = socks5::build_udp_header(&TargetAddr::Ip(dest));
dg.extend_from_slice(payload);
dg
}
fn acl_allow_all() -> DatagramAuthorizer {
Arc::new(|_, _, _| true)
}
fn client_datagrams(
socket: Arc<UdpSocket>,
client_ip: IpAddr,
lock: Arc<OnceLock<SocketAddr>>,
) -> ClientDatagrams {
ClientDatagrams::new(
socket,
client_ip,
lock,
Arc::new(ActivityClock::new()),
None,
Metrics::new(),
)
}
#[tokio::test]
async fn client_datagrams_recv_strips_header_and_locks_endpoint() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let relay_addr = relay.local_addr().unwrap();
let cd = client_datagrams(
relay,
IpAddr::from([127, 0, 0, 1]),
Arc::new(OnceLock::new()),
);
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr = client.local_addr().unwrap();
let dest: SocketAddr = "198.51.100.5:443".parse().unwrap();
client
.send_to(&socks_dg(dest, b"quic"), relay_addr)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
let (len, origin) = tokio::time::timeout(Duration::from_secs(2), cd.recv(&mut buf))
.await
.expect("a valid datagram should be delivered")
.unwrap();
assert_eq!(&buf[..len], b"quic", "the SOCKS header must be stripped");
assert_eq!(origin, dest, "the addressed origin is reported");
assert_eq!(
cd.client_endpoint(),
Some(client_addr),
"the first validated datagram locks the client endpoint"
);
}
#[tokio::test]
async fn client_datagrams_recv_rejects_foreign_source() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let relay_addr = relay.local_addr().unwrap();
let cd = client_datagrams(
relay,
IpAddr::from([127, 0, 0, 9]),
Arc::new(OnceLock::new()),
);
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
client
.send_to(
&socks_dg("198.51.100.5:443".parse().unwrap(), b"quic"),
relay_addr,
)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
assert!(
tokio::time::timeout(Duration::from_millis(300), cd.recv(&mut buf))
.await
.is_err(),
"a datagram from a foreign source must be dropped, not returned"
);
}
#[tokio::test]
async fn client_datagrams_recv_drops_fragment_keeps_good() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let relay_addr = relay.local_addr().unwrap();
let cd = client_datagrams(
relay,
IpAddr::from([127, 0, 0, 1]),
Arc::new(OnceLock::new()),
);
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let dest: SocketAddr = "198.51.100.5:443".parse().unwrap();
let mut fragmented = socks_dg(dest, b"frag");
fragmented[2] = 1; client.send_to(&fragmented, relay_addr).await.unwrap();
client
.send_to(&socks_dg(dest, b"good"), relay_addr)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
let (len, _) = tokio::time::timeout(Duration::from_secs(2), cd.recv(&mut buf))
.await
.expect("the unfragmented datagram should be delivered")
.unwrap();
assert_eq!(
&buf[..len],
b"good",
"a fragmented datagram must never be returned"
);
}
#[tokio::test]
async fn client_datagrams_recv_drops_domain_addressed() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let relay_addr = relay.local_addr().unwrap();
let cd = client_datagrams(
relay,
IpAddr::from([127, 0, 0, 1]),
Arc::new(OnceLock::new()),
);
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let mut domain = socks5::build_udp_header(&TargetAddr::Domain("example.com".into(), 443));
domain.extend_from_slice(b"dom");
client.send_to(&domain, relay_addr).await.unwrap();
client
.send_to(
&socks_dg("198.51.100.5:443".parse().unwrap(), b"ip"),
relay_addr,
)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
let (len, _) = tokio::time::timeout(Duration::from_secs(2), cd.recv(&mut buf))
.await
.expect("the IP-addressed datagram should be delivered")
.unwrap();
assert_eq!(
&buf[..len],
b"ip",
"a domain-addressed datagram is unsupported on takeover and must be dropped"
);
}
#[tokio::test]
async fn client_datagrams_send_reframes_reply_to_client() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let relay_addr = relay.local_addr().unwrap();
let cd = client_datagrams(
relay,
IpAddr::from([127, 0, 0, 1]),
Arc::new(OnceLock::new()),
);
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
client
.send_to(
&socks_dg("198.51.100.5:443".parse().unwrap(), b"hi"),
relay_addr,
)
.await
.unwrap();
let mut buf = vec![0u8; 2048];
cd.recv(&mut buf).await.unwrap();
let origin: SocketAddr = "203.0.113.1:443".parse().unwrap();
cd.send(origin, b"pong").await.unwrap();
let mut cbuf = vec![0u8; 2048];
let (n, from) = tokio::time::timeout(Duration::from_secs(2), client.recv_from(&mut cbuf))
.await
.expect("the client should receive the reply")
.unwrap();
assert_eq!(from, relay_addr);
let hdr = socks5::parse_udp_header(&cbuf[..n]).unwrap();
assert_eq!(
hdr.dest,
TargetAddr::Ip(origin),
"the reply carries the origin in its SOCKS header"
);
assert_eq!(&cbuf[hdr.payload_offset..n], b"pong");
}
#[tokio::test]
async fn client_datagrams_send_before_lock_is_noop() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let cd = client_datagrams(
relay,
IpAddr::from([127, 0, 0, 1]),
Arc::new(OnceLock::new()),
);
let would_be_target = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let from = would_be_target.local_addr().unwrap();
cd.send(from, b"pong").await.unwrap();
assert!(cd.client_endpoint().is_none());
let mut b = [0u8; 64];
assert!(
tokio::time::timeout(
Duration::from_millis(300),
would_be_target.recv_from(&mut b)
)
.await
.is_err(),
"send before the endpoint lock must not deliver anywhere, not even to `from`"
);
}
fn originator(
outbound: Arc<UdpSocket>,
policy: DnsPolicy,
acl: DatagramAuthorizer,
contacted: Arc<Mutex<ContactedRemotes>>,
) -> UpstreamOriginator {
UpstreamOriginator::new(
outbound,
false,
policy,
acl,
contacted,
Arc::new(ActivityClock::new()),
None,
Metrics::new(),
)
}
#[tokio::test]
async fn upstream_authorize_gates_on_acl() {
let outbound = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let acl: DatagramAuthorizer = Arc::new(|_h, _ip, port| port != 1);
let orig = originator(
outbound,
permissive_policy(),
acl,
Arc::new(Mutex::new(ContactedRemotes::new(true))),
);
let allowed: SocketAddr = "203.0.113.7:443".parse().unwrap();
assert_eq!(
orig.authorize(None, allowed).map(|t| t.dst()),
Some(allowed),
"an ACL-allowed destination yields a token"
);
assert!(
orig.authorize(None, "203.0.113.7:1".parse().unwrap())
.is_none(),
"an ACL-denied destination yields no token"
);
}
#[tokio::test]
async fn upstream_authorize_denies_dns_deny_category() {
let outbound = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let policy = DnsPolicy {
deny: vec![DnsDenyCategory::Loopback],
..permissive_policy()
};
let orig = originator(
outbound,
policy,
acl_allow_all(),
Arc::new(Mutex::new(ContactedRemotes::new(true))),
);
assert!(
orig.authorize(None, "127.0.0.1:443".parse().unwrap())
.is_none(),
"a loopback destination is denied by the DNS-deny policy"
);
assert!(
orig.authorize(None, "203.0.113.7:443".parse().unwrap())
.is_some(),
"a public destination passes"
);
}
#[tokio::test]
async fn upstream_send_to_reaches_origin_and_accepts_its_reply() {
let outbound = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let outbound_addr = outbound.local_addr().unwrap();
let origin = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let origin_addr = origin.local_addr().unwrap();
let contacted = Arc::new(Mutex::new(ContactedRemotes::new(true)));
let orig = originator(
outbound,
permissive_policy(),
acl_allow_all(),
contacted.clone(),
);
let target = orig
.authorize(None, origin_addr)
.expect("permissive authorize");
orig.send_to(&target, b"ping").await.unwrap();
let mut b = vec![0u8; 64];
let (n, _) = tokio::time::timeout(Duration::from_secs(2), origin.recv_from(&mut b))
.await
.expect("the origin should receive the datagram")
.unwrap();
assert_eq!(&b[..n], b"ping");
assert!(
contacted.lock().unwrap().contains(SocketAddr::new(
origin_addr.ip().to_canonical(),
origin_addr.port()
)),
"send_to records the destination as contacted"
);
origin.send_to(b"pong", outbound_addr).await.unwrap();
let (n, from) = tokio::time::timeout(Duration::from_secs(2), orig.recv(&mut b))
.await
.expect("a reply from a contacted origin should be delivered")
.unwrap();
assert_eq!(&b[..n], b"pong");
assert_eq!(from.ip(), origin_addr.ip().to_canonical());
}
#[tokio::test]
async fn upstream_recv_drops_uncontacted_reply() {
let outbound = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let outbound_addr = outbound.local_addr().unwrap();
let orig = originator(
outbound,
permissive_policy(),
acl_allow_all(),
Arc::new(Mutex::new(ContactedRemotes::new(true))),
);
let stranger = UdpSocket::bind("127.0.0.1:0").await.unwrap();
stranger.send_to(b"inject", outbound_addr).await.unwrap();
let mut b = vec![0u8; 64];
assert!(
tokio::time::timeout(Duration::from_millis(300), orig.recv(&mut b))
.await
.is_err(),
"an unsolicited reply from an uncontacted source must be dropped"
);
}
#[tokio::test]
async fn for_interceptor_constructors_drive_splice_association() {
let relay = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let relay_addr = relay.local_addr().unwrap();
let outbound = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let client = ClientDatagrams::for_interceptor(relay, IpAddr::from([127, 0, 0, 1]), None);
let upstream = UpstreamOriginator::for_interceptor(outbound, false, true, acl_allow_all());
let args = crate::plugin::AssociationArgs::for_interceptor(
client,
upstream,
Duration::from_secs(5),
);
let spliced = tokio::spawn(crate::plugin::splice_association(args));
let echo = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap());
let echo_addr = echo.local_addr().unwrap();
{
let echo = echo.clone();
tokio::spawn(async move {
let mut b = [0u8; 512];
while let Ok((n, peer)) = echo.recv_from(&mut b).await {
let _ = echo.send_to(&b[..n], peer).await;
}
});
}
let cli = UdpSocket::bind("127.0.0.1:0").await.unwrap();
cli.send_to(&socks_dg(echo_addr, b"ping"), relay_addr)
.await
.unwrap();
let mut buf = [0u8; 512];
let (n, _) = tokio::time::timeout(Duration::from_secs(2), cli.recv_from(&mut buf))
.await
.expect("the spliced association should relay the echo back")
.unwrap();
let hdr = socks5::parse_udp_header(&buf[..n]).unwrap();
assert_eq!(hdr.dest, TargetAddr::Ip(echo_addr));
assert_eq!(&buf[hdr.payload_offset..n], b"ping");
spliced.abort();
}
}