use super::relay::{UdpRelayState, UdpSocketRelay, UnspecifiedClientUdpAddressPolicy};
use crate::server::Error;
use rama_core::bytes::Bytes;
use rama_core::error::ErrorContext as _;
use rama_core::extensions::{Extensions, ExtensionsRef};
use rama_core::telemetry::tracing;
use rama_core::{Service, error::BoxError};
use rama_net::address::SocketAddress;
use rama_udp::UdpSocket;
use std::net::IpAddr;
#[cfg(feature = "dns")]
use super::MaybeDnsResolver;
#[expect(clippy::too_many_arguments)]
pub(super) trait UdpPacketProxy: Send + Sync + 'static {
fn proxy_udp_packets(
&self,
_extensions: Extensions,
client_address: SocketAddress,
north: UdpSocket,
north_read_buf_size: usize,
south: UdpSocket,
south_read_buf_size: usize,
#[cfg(feature = "dns")] dns_resolver: MaybeDnsResolver,
unspecified_client_udp_address_policy: UnspecifiedClientUdpAddressPolicy,
tcp_peer_ip: Option<IpAddr>,
) -> impl Future<Output = Result<(), Error>> + Send;
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct DirectUdpRelay;
impl UdpPacketProxy for DirectUdpRelay {
async fn proxy_udp_packets(
&self,
_extensions: Extensions,
client_address: SocketAddress,
north: UdpSocket,
north_read_buf_size: usize,
south: UdpSocket,
south_read_buf_size: usize,
#[cfg(feature = "dns")] dns_resolver: MaybeDnsResolver,
unspecified_client_udp_address_policy: UnspecifiedClientUdpAddressPolicy,
tcp_peer_ip: Option<IpAddr>,
) -> Result<(), Error> {
let relay = UdpSocketRelay::new(
client_address,
north,
north_read_buf_size,
south,
south_read_buf_size,
)
.with_unspecified_client_address_policy(unspecified_client_udp_address_policy, tcp_peer_ip);
#[cfg(feature = "dns")]
let relay = relay.with_dns_resolver(&_extensions, dns_resolver);
let mut relay = relay;
loop {
match relay.recv().await.map_err(Error::service)? {
Some(UdpRelayState::ReadNorth(server_address)) => {
tracing::trace!("relay: north -> south @ {server_address}");
relay
.send_to_south(None, server_address)
.await
.map_err(Error::service)?
}
Some(UdpRelayState::ReadSouth(server_address)) => {
tracing::trace!("relay: south @ {server_address} -> north");
relay
.send_to_north(None, server_address)
.await
.map_err(Error::service)?
}
None => {
tracing::trace!("ignore dropped packet: nothing to relay");
}
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum RelayDirection {
North,
South,
}
#[derive(Debug, Clone)]
pub struct RelayRequest {
pub direction: RelayDirection,
pub server_address: SocketAddress,
pub payload: Bytes,
pub extensions: Extensions,
}
impl ExtensionsRef for RelayRequest {
fn extensions(&self) -> &Extensions {
&self.extensions
}
}
#[derive(Debug, Clone)]
pub struct RelayResponse {
pub maybe_payload: Option<Bytes>,
pub extensions: Extensions,
}
impl ExtensionsRef for RelayResponse {
fn extensions(&self) -> &Extensions {
&self.extensions
}
}
#[derive(Debug, Clone)]
pub struct AsyncUdpInspector<S>(pub(super) S);
impl<S> UdpPacketProxy for AsyncUdpInspector<S>
where
S: Service<RelayRequest, Output = RelayResponse, Error: Into<BoxError>>,
{
async fn proxy_udp_packets(
&self,
mut extensions: Extensions,
client_address: SocketAddress,
north: UdpSocket,
north_read_buf_size: usize,
south: UdpSocket,
south_read_buf_size: usize,
#[cfg(feature = "dns")] dns_resolver: MaybeDnsResolver,
unspecified_client_udp_address_policy: UnspecifiedClientUdpAddressPolicy,
tcp_peer_ip: Option<IpAddr>,
) -> Result<(), Error> {
let relay = UdpSocketRelay::new(
client_address,
north,
north_read_buf_size,
south,
south_read_buf_size,
)
.with_unspecified_client_address_policy(unspecified_client_udp_address_policy, tcp_peer_ip);
#[cfg(feature = "dns")]
let relay = relay.with_dns_resolver(&extensions, dns_resolver);
let mut relay = relay;
loop {
match relay.recv().await.map_err(Error::service)? {
Some(UdpRelayState::ReadNorth(server_address)) => {
tracing::trace!("relay request: north -> south @ {server_address}");
let request = RelayRequest {
direction: RelayDirection::South,
server_address,
payload: Bytes::copy_from_slice(relay.north_read_buf_slice()),
extensions,
};
let result = self
.0
.serve(request)
.await
.into_box_error()
.inspect_err(|err| {
tracing::debug!(
"relay request: south @ {server_address} -> north: failed: {err:?}"
);
})
.map_err(Error::service)?;
let maybe_payload;
RelayResponse {
extensions,
maybe_payload,
} = result;
match maybe_payload {
Some(payload) => relay
.send_to_south(Some(payload), server_address)
.await
.map_err(Error::service)?,
None => {
tracing::trace!(
"block request: north -> south @ {server_address}: inspector blocked"
);
}
}
}
Some(UdpRelayState::ReadSouth(server_address)) => {
tracing::trace!("relay request: south @ {server_address} -> north");
let request = RelayRequest {
direction: RelayDirection::North,
server_address,
payload: Bytes::copy_from_slice(relay.south_read_buf_slice()),
extensions,
};
let result = self
.0
.serve(request)
.await
.into_box_error()
.inspect_err(|err| {
tracing::debug!(
"relay request: north -> south @ {server_address}: failed: {err:?}"
);
})
.map_err(Error::service)?;
let maybe_payload;
RelayResponse {
extensions,
maybe_payload,
} = result;
match maybe_payload {
Some(payload) => relay
.send_to_north(Some(payload), server_address)
.await
.map_err(Error::service)?,
None => {
tracing::trace!(
"block request: south @ {server_address} -> north: inspector blocked"
);
}
}
}
None => {
tracing::trace!("ignore dropped packet: nothing to inspect or relay");
}
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum UdpInspectAction {
Forward,
Block,
Modify(Bytes),
}
pub trait UdpInspector: Send + Sync + 'static {
type Error: Into<BoxError> + Send + 'static;
fn inspect_packet(
&self,
direction: RelayDirection,
server_address: SocketAddress,
payload: &[u8],
) -> Result<UdpInspectAction, Self::Error>;
}
impl<F, E> UdpInspector for F
where
F: Fn(RelayDirection, SocketAddress, &[u8]) -> Result<UdpInspectAction, E>
+ Send
+ Sync
+ 'static,
E: Into<BoxError> + Send + 'static,
{
type Error = E;
fn inspect_packet(
&self,
direction: RelayDirection,
server_address: SocketAddress,
payload: &[u8],
) -> Result<UdpInspectAction, Self::Error> {
(self)(direction, server_address, payload)
}
}
#[derive(Debug, Clone)]
pub struct SyncUdpInspector<S>(pub(super) S);
impl<S> UdpPacketProxy for SyncUdpInspector<S>
where
S: UdpInspector,
{
async fn proxy_udp_packets(
&self,
_extensions: Extensions,
client_address: SocketAddress,
north: UdpSocket,
north_read_buf_size: usize,
south: UdpSocket,
south_read_buf_size: usize,
#[cfg(feature = "dns")] dns_resolver: MaybeDnsResolver,
unspecified_client_udp_address_policy: UnspecifiedClientUdpAddressPolicy,
tcp_peer_ip: Option<IpAddr>,
) -> Result<(), Error> {
let relay = UdpSocketRelay::new(
client_address,
north,
north_read_buf_size,
south,
south_read_buf_size,
)
.with_unspecified_client_address_policy(unspecified_client_udp_address_policy, tcp_peer_ip);
#[cfg(feature = "dns")]
let relay = relay.with_dns_resolver(&_extensions, dns_resolver);
let mut relay = relay;
loop {
match relay.recv().await.map_err(Error::service)? {
Some(UdpRelayState::ReadNorth(server_address)) => {
tracing::trace!("relay request: north -> south @ {server_address}");
let action = self
.0
.inspect_packet(
RelayDirection::South,
server_address,
relay.north_read_buf_slice(),
)
.into_box_error()
.inspect_err(|err| {
tracing::debug!(
"relay request: north -> south @ {server_address}: failed: {err:?}"
);
})
.map_err(Error::service)?;
match action {
UdpInspectAction::Forward => {
tracing::trace!(
"relay request: north -> south @ {server_address}: forward"
);
relay
.send_to_south(None, server_address)
.await
.map_err(Error::service)?;
}
UdpInspectAction::Block => {
tracing::trace!(
"block request: north -> south @ {server_address}: inspector blocked"
);
}
UdpInspectAction::Modify(bytes) => {
tracing::trace!(
"relay request: north -> south @ {server_address}: forward modified bytes (len = {})",
bytes.len()
);
relay
.send_to_south(Some(bytes), server_address)
.await
.map_err(Error::service)?;
}
}
}
Some(UdpRelayState::ReadSouth(server_address)) => {
tracing::trace!("relay request: south @ {server_address} -> north");
let action = self
.0
.inspect_packet(
RelayDirection::North,
server_address,
relay.south_read_buf_slice(),
)
.into_box_error()
.inspect_err(|err| {
tracing::debug!(
"relay request: south @ {server_address} -> north: failed: {err:?}"
);
})
.map_err(Error::service)?;
match action {
UdpInspectAction::Forward => {
tracing::trace!(
"relay request: south @ {server_address} -> north: forward"
);
relay
.send_to_north(None, server_address)
.await
.map_err(Error::service)?;
}
UdpInspectAction::Block => {
tracing::trace!(
"block request: south @ {server_address} -> north: inspector blocked"
);
}
UdpInspectAction::Modify(bytes) => {
tracing::trace!(
"relay request: south @ {server_address} -> north: forward modified bytes (len = {})",
bytes.len()
);
relay
.send_to_north(Some(bytes), server_address)
.await
.map_err(Error::service)?;
}
}
}
None => {
tracing::trace!("ignore dropped packet: nothing to inspect or relay");
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use rama_core::error::BoxError;
use rama_core::service::service_fn;
#[tokio::test]
async fn test_async_inspector_south_to_north_routes_to_client() {
let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let client_addr: std::net::SocketAddr = client_socket.local_addr().unwrap();
let client_socket_addr: SocketAddress = client_addr.into();
let north = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let south = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let south_addr = south.local_addr().unwrap();
let server = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let inspector = AsyncUdpInspector(service_fn(async move |req: RelayRequest| {
Ok::<_, BoxError>(RelayResponse {
maybe_payload: Some(req.payload),
extensions: req.extensions,
})
}));
let payload = b"south_to_north_regression";
server.send_to(payload, south_addr).await.unwrap();
let mut buf = vec![0u8; 4096];
let inspect_fut = inspector.proxy_udp_packets(
Extensions::new(),
client_socket_addr,
north,
4096,
south,
4096,
#[cfg(feature = "dns")]
Default::default(),
UnspecifiedClientUdpAddressPolicy::default(),
None,
);
tokio::select! {
_ = inspect_fut => {
panic!("inspector exited unexpectedly before packet arrived");
}
result = tokio::time::timeout(
std::time::Duration::from_millis(500),
client_socket.recv_from(&mut buf),
) => {
let (n, _) = result
.expect("timed out: south→north packet did not arrive at client (north side)")
.unwrap();
let received = &buf[..n];
assert!(
received.ends_with(payload),
"south→north: payload must arrive at client (north side)"
);
}
}
}
}