use rama_core::bytes::{BufMut, BytesMut};
use rama_core::error::BoxErrorExt as _;
use rama_core::error::{BoxError, ErrorContext as _, ErrorExt as _};
use rama_core::futures::Sink;
use rama_core::futures::Stream;
use rama_core::stream::codec::{Decoder, Encoder};
use rama_core::telemetry::tracing;
use rama_net::address::HostWithPort;
use rama_net::address::SocketAddress;
use rama_udp::{UdpSocket, bind_udp_with_address};
use rama_utils::octets::kib;
use std::pin::Pin;
use std::task::{Context, Poll, ready};
use std::{fmt, io, net::SocketAddr};
use tokio::io::ReadBuf;
use crate::proto::{Command, ProtocolVersion, ReplyKind, client::Request, server, udp::UdpHeader};
use super::core::HandshakeError;
pub struct UdpSocketRelayBinder<S> {
stream: S,
}
impl<S: rama_core::io::Io + Unpin> UdpSocketRelayBinder<S> {
pub(crate) fn new(stream: S) -> Self {
Self { stream }
}
pub async fn bind_address(
mut self,
address: impl TryInto<SocketAddress, Error: Into<BoxError>>,
) -> Result<UdpSocketRelay<S>, HandshakeError> {
let socket = bind_udp_with_address(address).await.map_err(|err| {
HandshakeError::other(err).with_context("bind udp socket ready for sending")
})?;
let socket_addr = socket.local_addr().map_err(|err| {
HandshakeError::other(err).with_context("get local address from udp sender socker")
})?;
let request = Request {
version: ProtocolVersion::Socks5,
command: Command::UdpAssociate,
destination: socket_addr.into(),
};
request.write_to(&mut self.stream).await.map_err(|err| {
HandshakeError::io(err).with_context("write client request: UDP Associate")
})?;
tracing::trace!(
network.local.address = %socket_addr.ip(),
network.local.port = %socket_addr.port(),
"socks5 client: udp associate handshake initiated"
);
let server_reply = server::Reply::read_from(&mut self.stream)
.await
.map_err(|err| HandshakeError::protocol(err).with_context("read server reply"))?;
if server_reply.reply != ReplyKind::Succeeded {
return Err(HandshakeError::reply_kind(server_reply.reply)
.with_context("server responded with non-success reply"));
}
let HostWithPort { host, port } = server_reply.bind_address;
let bind_address: SocketAddress = match host.try_as_ip() {
Ok(ip) => (ip, port).into(),
Err(_) => {
return Err(
HandshakeError::reply_kind(ReplyKind::AddressTypeNotSupported).with_context(
"server responded with non-IP address: incompatible for udp bind",
),
);
}
};
socket
.connect(bind_address.into_std())
.await
.map_err(|err| {
HandshakeError::other(err).with_context("connect to socks5 udp association socket")
})?;
tracing::trace!(
network.local.address = %socket_addr.ip(),
network.local.port = %socket_addr.port(),
"socks5 client: socks5 server ready to bind at {bind_address} for udp purposes",
);
Ok(UdpSocketRelay {
stream: self.stream,
socket,
write_buffer: BytesMut::with_capacity(1024),
})
}
}
pub struct UdpSocketRelay<S> {
stream: S,
socket: UdpSocket,
write_buffer: BytesMut,
}
impl<S: fmt::Debug> fmt::Debug for UdpSocketRelay<S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UdpSocketRelay")
.field("stream", &self.stream)
.field("socket", &self.socket)
.field("write_buffer", &self.write_buffer)
.finish()
}
}
impl<S: rama_core::io::Io + Unpin> UdpSocketRelay<S> {
#[inline]
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.socket.local_addr()
}
#[inline]
pub fn peer_addr(&self) -> io::Result<SocketAddr> {
self.socket.peer_addr()
}
pub async fn send_to<A: TryInto<SocketAddress, Error: Into<BoxError>>>(
&mut self,
b: &[u8],
addr: A,
) -> Result<usize, BoxError> {
let socket_addr: SocketAddress = addr.try_into().into_box_error()?;
let header = UdpHeader {
fragment_number: 0,
destination: socket_addr.into(),
};
self.write_buffer.truncate(0);
header.write_to_buf(&mut self.write_buffer)?;
self.write_buffer.extend_from_slice(b);
Ok(self.socket.send(&self.write_buffer[..]).await?)
}
pub fn poll_send_to<A: TryInto<SocketAddress, Error: Into<BoxError>>>(
&mut self,
cx: &mut Context<'_>,
b: &[u8],
addr: A,
) -> Poll<Result<usize, BoxError>> {
let socket_addr: SocketAddress = addr.try_into().into_box_error()?;
let header = UdpHeader {
fragment_number: 0,
destination: socket_addr.into(),
};
self.write_buffer.truncate(0);
header.write_to_buf(&mut self.write_buffer)?;
self.write_buffer.extend_from_slice(b);
self.socket
.poll_send(cx, &self.write_buffer[..])
.map_err(Into::into)
.map_ok(|n| n - header.serialized_len())
}
pub async fn recv_from(&mut self, buf: &mut [u8]) -> Result<(usize, SocketAddress), BoxError> {
let n = self.socket.recv(buf).await?;
let header = UdpHeader::read_from(&mut &buf[..n]).await?;
let (header_offset, from) = validate_udp_header(header)?;
buf.copy_within(header_offset.., 0);
Ok((n - header_offset, from))
}
pub fn poll_recv_from(
&mut self,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<Result<SocketAddress, BoxError>> {
if let Err(err) = ready!(self.socket.poll_recv(cx, buf)) {
return Poll::Ready(Err(err.into()));
}
let header = UdpHeader::read_from_sync(&mut buf.filled())?;
let (header_offset, from) = validate_udp_header(header)?;
let filled = buf.filled_mut();
let len = filled.len();
assert!(len > header_offset);
filled.copy_within(header_offset.., 0);
buf.set_filled(len - header_offset);
Poll::Ready(Ok(from))
}
pub fn into_framed<C>(self, codec: C) -> UdpFramedRelay<C, S> {
UdpFramedRelay::new(self, codec)
}
}
fn validate_udp_header(header: UdpHeader) -> Result<(usize, SocketAddress), BoxError> {
if header.fragment_number != 0 {
return Err(BoxError::from_static_str(
"UdpSocketRelay: fragment number != 0 is not supported",
)
.context_field("fragment_number", header.fragment_number));
}
let header_offset = header.serialized_len();
let HostWithPort { host, port } = header.destination;
let from: SocketAddress = match host.try_as_ip() {
Ok(ip) => (ip, port).into(),
Err(_) => {
return Err(BoxError::from_static_str(
"server responded with non-IP host: incompatible for udp bind",
)
.context_field("host", host.to_string()));
}
};
Ok((header_offset, from))
}
impl<S: Send + Sync + 'static> rama_net::stream::Socket for UdpSocketRelay<S> {
fn local_addr(&self) -> io::Result<SocketAddress> {
self.socket.local_addr().map(Into::into)
}
fn peer_addr(&self) -> io::Result<SocketAddress> {
self.socket.peer_addr().map(Into::into)
}
}
#[must_use = "sinks do nothing unless polled"]
pub struct UdpFramedRelay<C, S> {
relay_socket: UdpSocketRelay<S>,
codec: C,
rd: BytesMut,
wr: BytesMut,
out_addr: SocketAddress,
flushed: bool,
current_addr: Option<SocketAddress>,
}
impl<C, S: rama_core::io::Io + Unpin> UdpFramedRelay<C, S> {
#[inline]
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.relay_socket.local_addr()
}
#[inline]
pub fn peer_addr(&self) -> io::Result<SocketAddr> {
self.relay_socket.peer_addr()
}
}
impl<C: fmt::Debug, S: fmt::Debug> fmt::Debug for UdpFramedRelay<C, S> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UdpFramedRelay")
.field("relay_socket", &self.relay_socket)
.field("codec", &self.relay_socket)
.field("rd", &self.rd)
.field("wr", &self.wr)
.field("out_addr", &self.out_addr)
.field("flushed", &self.flushed)
.field("current_addr", &self.current_addr)
.finish()
}
}
impl<C: Send + Sync + 'static, S: Send + Sync + 'static> rama_net::stream::Socket
for UdpFramedRelay<C, S>
{
fn local_addr(&self) -> io::Result<SocketAddress> {
self.relay_socket.local_addr()
}
fn peer_addr(&self) -> io::Result<SocketAddress> {
self.relay_socket.peer_addr()
}
}
const INITIAL_RD_CAPACITY: usize = kib(64);
const INITIAL_WR_CAPACITY: usize = kib(8);
impl<C, S: Unpin> Unpin for UdpFramedRelay<C, S> {}
impl<C, S> Stream for UdpFramedRelay<C, S>
where
C: Decoder<Error: Into<BoxError>>,
S: rama_core::io::Io + Unpin,
{
type Item = Result<(C::Item, SocketAddress), BoxError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let pin = self.get_mut();
pin.rd.reserve(INITIAL_RD_CAPACITY);
loop {
if let Some(current_addr) = pin.current_addr {
if let Some(frame) = pin.codec.decode_eof(&mut pin.rd).into_box_error()? {
return Poll::Ready(Some(Ok((frame, current_addr))));
}
pin.current_addr = None;
pin.rd.clear();
}
let addr = {
let buf = unsafe { pin.rd.chunk_mut().as_uninit_slice_mut() };
let mut read = ReadBuf::uninit(buf);
let ptr = read.filled().as_ptr();
let res = ready!(pin.relay_socket.poll_recv_from(cx, &mut read));
assert_eq!(ptr, read.filled().as_ptr());
let addr = res?;
let filled = read.filled().len();
unsafe { pin.rd.advance_mut(filled) };
addr
};
pin.current_addr = Some(addr);
}
}
}
impl<I, C, S> Sink<(I, SocketAddr)> for UdpFramedRelay<C, S>
where
C: Encoder<I, Error: Into<BoxError>>,
S: rama_core::io::Io + Unpin,
{
type Error = BoxError;
fn poll_ready(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if !self.flushed {
match self.poll_flush(cx)? {
Poll::Ready(()) => {}
Poll::Pending => return Poll::Pending,
}
}
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, item: (I, SocketAddr)) -> Result<(), Self::Error> {
let (frame, out_addr) = item;
let pin = self.get_mut();
pin.codec.encode(frame, &mut pin.wr).into_box_error()?;
pin.out_addr = out_addr.into();
pin.flushed = false;
Ok(())
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
if self.flushed {
return Poll::Ready(Ok(()));
}
let Self {
ref mut relay_socket,
ref mut out_addr,
ref mut wr,
..
} = *self;
let n = ready!(relay_socket.poll_send_to(cx, wr, *out_addr))?;
let wr_n = self.wr.len();
let wrote_all = n == wr_n;
self.wr.clear();
self.flushed = true;
let res = if wrote_all {
Ok(())
} else {
tracing::debug!(
"failed to write entire datagram to socket: len = {n}; wr len = {wr_n}"
);
Err(io::Error::other("failed to write entire datagram to socket").into())
};
Poll::Ready(res)
}
fn poll_close(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
ready!(self.poll_flush(cx))?;
Poll::Ready(Ok(()))
}
}
impl<C, S> UdpFramedRelay<C, S> {
fn new(relay_socket: UdpSocketRelay<S>, codec: C) -> Self {
Self {
relay_socket,
codec,
out_addr: SocketAddress::default_ipv4(0),
rd: BytesMut::with_capacity(INITIAL_RD_CAPACITY),
wr: BytesMut::with_capacity(INITIAL_WR_CAPACITY),
flushed: true,
current_addr: None,
}
}
pub fn codec(&self) -> &C {
&self.codec
}
pub fn codec_mut(&mut self) -> &mut C {
&mut self.codec
}
}