use std::time::Duration;
use rama_core::{
Service, combinators::Either, error::BoxError, extensions::ExtensionsRef, io::Io,
layer::timeout::DefaultTimeout, telemetry::tracing,
};
use rama_net::{
address::{HostWithPort, SocketAddress},
socket::SocketService,
stream::SocketInfo,
};
use rama_udp::{UdpSocket, bind_udp_with_address};
use rama_utils::macros::generate_set_and_with;
#[cfg(feature = "dns")]
use ::rama_dns::client::resolver::{BoxDnsAddressResolver, DnsAddressResolver};
use super::Error;
use crate::proto::{ReplyKind, server::Reply};
mod inspect;
use inspect::UdpPacketProxy;
pub use inspect::{
AsyncUdpInspector, DirectUdpRelay, RelayDirection, RelayRequest, RelayResponse,
SyncUdpInspector, UdpInspectAction, UdpInspector,
};
mod relay;
pub use relay::UnspecifiedClientUdpAddressPolicy;
#[cfg(feature = "dns")]
type MaybeDnsResolver = Option<BoxDnsAddressResolver>;
pub trait Socks5UdpAssociator<S>: Socks5UdpAssociatorSeal<S> {}
impl<S, C> Socks5UdpAssociator<S> for C where C: Socks5UdpAssociatorSeal<S> {}
pub trait Socks5UdpAssociatorSeal<S>: Send + Sync + 'static {
fn accept_udp_associate(
&self,
stream: S,
destination: HostWithPort,
) -> impl Future<Output = Result<(), Error>> + Send + '_
where
S: Io + Unpin;
}
impl<S> Socks5UdpAssociatorSeal<S> for ()
where
S: Io + Unpin,
{
async fn accept_udp_associate(
&self,
mut stream: S,
destination: HostWithPort,
) -> Result<(), Error> {
tracing::debug!(
"socks5 server w/ destination {destination}: abort: command not supported: UDP Associate",
);
Reply::error_reply(ReplyKind::CommandNotSupported)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err)
.with_context("write server reply: command not supported (udp associate)")
})?;
Err(Error::aborted("command not supported: UDP Associate"))
}
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct DefaultUdpBinder;
impl Service<SocketAddress> for DefaultUdpBinder {
type Output = UdpSocket;
type Error = BoxError;
async fn serve(&self, addr: SocketAddress) -> Result<Self::Output, Self::Error> {
let socket = bind_udp_with_address(addr).await?;
Ok(socket)
}
}
pub type DefaultUdpRelay = UdpRelay<DefaultTimeout<DefaultUdpBinder>, DirectUdpRelay>;
#[derive(Debug, Clone)]
pub struct UdpRelay<B, I> {
binder: B,
inspector: I,
#[cfg(feature = "dns")]
dns_resolver: MaybeDnsResolver,
bind_north_address: SocketAddress,
bind_south_address: SocketAddress,
north_buffer_size: usize,
south_buffer_size: usize,
relay_timeout: Option<Duration>,
unspecified_client_udp_address_policy: UnspecifiedClientUdpAddressPolicy,
}
impl<B> UdpRelay<B, DirectUdpRelay> {
pub fn new(binder: B) -> Self {
Self {
binder,
inspector: DirectUdpRelay::default(),
#[cfg(feature = "dns")]
dns_resolver: Default::default(),
bind_north_address: SocketAddress::default_ipv4(0),
bind_south_address: SocketAddress::default_ipv4(0),
north_buffer_size: 4096,
south_buffer_size: 4096,
relay_timeout: None,
unspecified_client_udp_address_policy: UnspecifiedClientUdpAddressPolicy::default(),
}
}
pub fn with_sync_inspector<T>(self, inspector: T) -> UdpRelay<B, SyncUdpInspector<T>> {
UdpRelay {
binder: self.binder,
inspector: SyncUdpInspector(inspector),
#[cfg(feature = "dns")]
dns_resolver: self.dns_resolver,
bind_north_address: self.bind_north_address,
bind_south_address: self.bind_south_address,
north_buffer_size: self.north_buffer_size,
south_buffer_size: self.south_buffer_size,
relay_timeout: self.relay_timeout,
unspecified_client_udp_address_policy: self.unspecified_client_udp_address_policy,
}
}
pub fn with_async_inspector<T>(self, inspector: T) -> UdpRelay<B, AsyncUdpInspector<T>> {
UdpRelay {
binder: self.binder,
inspector: AsyncUdpInspector(inspector),
#[cfg(feature = "dns")]
dns_resolver: self.dns_resolver,
bind_north_address: self.bind_north_address,
bind_south_address: self.bind_south_address,
north_buffer_size: self.north_buffer_size,
south_buffer_size: self.south_buffer_size,
relay_timeout: self.relay_timeout,
unspecified_client_udp_address_policy: self.unspecified_client_udp_address_policy,
}
}
}
impl<B, I> UdpRelay<B, I> {
pub fn with_binder<T>(self, binder: T) -> UdpRelay<T, I> {
UdpRelay {
binder,
inspector: self.inspector,
#[cfg(feature = "dns")]
dns_resolver: self.dns_resolver,
bind_north_address: self.bind_north_address,
bind_south_address: self.bind_south_address,
north_buffer_size: self.north_buffer_size,
south_buffer_size: self.south_buffer_size,
relay_timeout: self.relay_timeout,
unspecified_client_udp_address_policy: self.unspecified_client_udp_address_policy,
}
}
generate_set_and_with! {
pub fn bind_address(mut self, address: impl Into<SocketAddress>) -> Self {
let address = address.into();
self.bind_north_address = address;
self.bind_south_address = address;
self
}
}
generate_set_and_with! {
pub fn bind_north_address(mut self, address: impl Into<SocketAddress>) -> Self {
self.bind_north_address = address.into();
self
}
}
generate_set_and_with! {
pub fn bind_south_address(mut self, address: impl Into<SocketAddress>) -> Self {
self.bind_south_address = address.into();
self
}
}
generate_set_and_with! {
pub fn buffer_size_south(mut self, n: usize) -> Self {
self.south_buffer_size = n;
self
}
}
generate_set_and_with! {
pub fn buffer_size_north(mut self, n: usize) -> Self {
self.north_buffer_size = n;
self
}
}
generate_set_and_with! {
pub fn buffer_size(mut self, n: usize) -> Self {
self.north_buffer_size = n;
self.south_buffer_size = n;
self
}
}
generate_set_and_with! {
pub fn relay_timeout(mut self, timeout: Option<Duration>) -> Self {
self.relay_timeout = timeout;
self
}
}
generate_set_and_with! {
pub fn unspecified_client_udp_address_policy(
mut self,
policy: UnspecifiedClientUdpAddressPolicy,
) -> Self {
self.unspecified_client_udp_address_policy = policy;
self
}
}
}
#[cfg(feature = "dns")]
impl<B, I> UdpRelay<B, I> {
generate_set_and_with! {
pub fn default_dns_resolver(mut self) -> Self {
self.dns_resolver = Some(::rama_dns::client::GlobalDnsResolver::new().into_box_dns_address_resolver());
self
}
}
generate_set_and_with! {
pub fn dns_resolver(mut self, resolver: Option<BoxDnsAddressResolver>) -> Self {
self.dns_resolver = resolver;
self
}
}
#[must_use]
pub fn with_dns_address_resolver(mut self, resolver: impl DnsAddressResolver) -> Self {
self.dns_resolver = Some(resolver.into_box_dns_address_resolver());
self
}
pub fn set_dns_address_resolver(&mut self, resolver: impl DnsAddressResolver) -> &mut Self {
self.dns_resolver = Some(resolver.into_box_dns_address_resolver());
self
}
}
impl Default for DefaultUdpRelay {
fn default() -> Self {
let relay = Self::new(DefaultTimeout::new(
DefaultUdpBinder::default(),
Duration::from_secs(30),
))
.with_relay_timeout(Duration::from_secs(300));
#[cfg(feature = "dns")]
let relay = relay.with_default_dns_resolver();
relay
}
}
impl<B, I, S> Socks5UdpAssociatorSeal<S> for UdpRelay<B, I>
where
B: SocketService<Socket = UdpSocket>,
I: UdpPacketProxy,
S: Io + Unpin + ExtensionsRef,
{
async fn accept_udp_associate(
&self,
mut stream: S,
destination: HostWithPort,
) -> Result<(), Error> {
tracing::trace!(
"socks5 server w/ destination {destination}: udp associate: try to bind incoming socket to destination {destination}",
);
let extensions = stream.extensions().clone();
let HostWithPort {
host: dest_host,
port: dest_port,
} = destination;
let Ok(dest_addr) = dest_host.try_as_ip() else {
tracing::debug!(
"udp associate command does not accept non-IP host {dest_host} as bind address",
);
let reply_kind = ReplyKind::AddressTypeNotSupported;
Reply::error_reply(reply_kind)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: udp relay failed")
})?;
return Err(Error::aborted("udp relay failed").with_context(reply_kind));
};
let client_address = SocketAddress::new(dest_addr, dest_port);
let tcp_peer_ip = extensions
.get_ref::<SocketInfo>()
.map(|info| info.peer_addr().ip_addr);
if client_address.ip_addr.is_unspecified()
&& self.unspecified_client_udp_address_policy
== UnspecifiedClientUdpAddressPolicy::PinToTcpPeerIp
&& tcp_peer_ip.is_none()
{
tracing::warn!(
"socks5 udp associate: PinToTcpPeerIp cannot enforce IP filtering \
(no SocketInfo / TCP peer IP available); degrading to first-packet pinning",
);
}
let socket_north = match self
.binder
.bind_socket_with_address(self.bind_north_address)
.await
{
Ok(twin) => twin,
Err(err) => {
let err = err.into();
tracing::debug!("udp north socket bind failed: {err:?}",);
let reply_kind = ReplyKind::GeneralServerFailure;
Reply::error_reply(reply_kind)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err)
.with_context("write server reply: udp north socket bind failed")
})?;
return Err(Error::aborted("udp north socket bind failed")
.with_context(reply_kind)
.with_source(err));
}
};
let socket_north_address = match socket_north.local_addr() {
Ok(addr) => addr,
Err(err) => {
tracing::debug!("retrieve local addr of north (udp) socket failed: {err:?}");
let reply_kind = ReplyKind::GeneralServerFailure;
Reply::error_reply(reply_kind)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err)
.with_context("write server reply: prepare udp receive socket failed")
})?;
return Err(
Error::aborted("prepare udp receive socket failed").with_context(reply_kind)
);
}
};
let socket_south = match self
.binder
.bind_socket_with_address(self.bind_south_address)
.await
{
Ok(twin) => twin,
Err(err) => {
let err = err.into();
tracing::debug!("udp south socket bind failed: {err:?}",);
let reply_kind = ReplyKind::GeneralServerFailure;
Reply::error_reply(reply_kind)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err)
.with_context("write server reply: udp south socket bind failed")
})?;
return Err(Error::aborted("udp south socket bind failed")
.with_context(reply_kind)
.with_source(err));
}
};
Reply::new(socket_north_address)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err)
.with_context("write server reply: udp associate: north+south sockets ready")
})?;
let mut empty = tokio::io::empty();
let mut drop_stream_fut = std::pin::pin!(tokio::io::copy(&mut stream, &mut empty));
let mut timeout_fut = std::pin::pin!(match self.relay_timeout {
Some(timeout) => Either::A(tokio::time::sleep(timeout)),
None => Either::B(std::future::pending::<()>()),
});
#[cfg(feature = "dns")]
let udp_relay = self.inspector.proxy_udp_packets(
extensions,
client_address,
socket_north,
self.north_buffer_size,
socket_south,
self.south_buffer_size,
self.dns_resolver.clone(),
self.unspecified_client_udp_address_policy,
tcp_peer_ip,
);
#[cfg(not(feature = "dns"))]
let udp_relay = self.inspector.proxy_udp_packets(
extensions,
client_address,
socket_north,
self.north_buffer_size,
socket_south,
self.south_buffer_size,
self.unspecified_client_udp_address_policy,
tcp_peer_ip,
);
tokio::select! {
_ = &mut drop_stream_fut => {
tracing::trace!(
network.peer.address = %client_address.ip_addr,
network.peer.port = %client_address.port,
"socks5 server: udp associate: tcp stream dropped from client: drop relay",
);
}
_ = &mut timeout_fut => {
tracing::debug!(
network.peer.address = %client_address.ip_addr,
network.peer.port = %client_address.port,
"socks5 server: udp associate: timeout reached: drop relay",
);
return Err(Error::io(std::io::Error::new(std::io::ErrorKind::TimedOut, "relay timeout reached")));
}
Err(err) = udp_relay => {
tracing::debug!(
network.peer.address = %client_address.ip_addr,
network.peer.port = %client_address.port,
"socks5 server: udp associate: udp relay: exit with an error",
);
return Err(err);
}
}
tracing::trace!(
network.peer.address = %client_address.ip_addr,
network.peer.port = %client_address.port,
"socks5 server: udp associate: udp relay: done",);
Ok(())
}
}
#[cfg(test)]
pub(crate) use test::MockUdpAssociator;
#[cfg(test)]
mod test;