use std::{io, time::Duration};
use rama_core::io::BridgeIo;
use rama_core::rt::Executor;
use rama_core::telemetry::tracing::{self, Instrument};
use rama_core::{Service, error::BoxError, io::Io, layer::timeout::DefaultTimeout};
use rama_net::address::HostWithPort;
use rama_net::{address::SocketAddress, proxy::IoForwardService, socket::SocketService};
use rama_tcp::{TcpStream, server::TcpListener};
use rama_utils::macros::generate_set_and_with;
use super::Error;
use crate::proto::{ReplyKind, server::Reply};
pub trait Socks5Binder<S>: Socks5BinderSeal<S> {}
impl<S, C> Socks5Binder<S> for C where C: Socks5BinderSeal<S> {}
pub trait Socks5BinderSeal<S>: Send + Sync + 'static {
fn accept_bind(
&self,
stream: S,
destination: HostWithPort,
) -> impl Future<Output = Result<(), Error>> + Send + '_;
}
impl<S> Socks5BinderSeal<S> for ()
where
S: Io + Unpin,
{
async fn accept_bind(&self, mut stream: S, destination: HostWithPort) -> Result<(), Error> {
tracing::debug!(
server.address = %destination.host,
server.port = %destination.port,
"socks5 server: abort: command not supported: Bind",
);
Reply::error_reply(ReplyKind::CommandNotSupported)
.write_to(&mut stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: command not supported (bind)")
})?;
Err(Error::aborted("command not supported: Bind"))
}
}
pub type DefaultBinder = Binder<DefaultTimeout<DefaultAcceptorFactory>, IoForwardService>;
#[derive(Debug, Clone)]
pub struct Binder<A, S> {
acceptor: A,
service: S,
bind_address: Option<SocketAddress>,
accept_timeout: Option<Duration>,
}
impl<A, S> Binder<A, S> {
pub fn new(acceptor: A, service: S) -> Self {
Self {
acceptor,
service,
bind_address: None,
accept_timeout: None,
}
}
}
impl<A, S> Binder<A, S> {
pub fn with_acceptor<T>(self, acceptor: T) -> Binder<T, S> {
Binder {
acceptor,
service: self.service,
bind_address: self.bind_address,
accept_timeout: self.accept_timeout,
}
}
pub fn with_service<T>(self, service: T) -> Binder<A, T> {
Binder {
acceptor: self.acceptor,
service,
bind_address: self.bind_address,
accept_timeout: self.accept_timeout,
}
}
generate_set_and_with! {
pub fn bind_address(mut self, addr: impl Into<SocketAddress>) -> Self {
self.bind_address = Some(addr.into());
self
}
}
generate_set_and_with! {
pub fn default_bind_address(mut self) -> Self {
self.bind_address = Some(SocketAddress::default_ipv4(0));
self
}
}
generate_set_and_with! {
pub fn accept_timeout(mut self, timeout: Option<Duration>) -> Self {
self.accept_timeout = timeout;
self
}
}
}
#[derive(Debug, Clone, Default)]
pub struct DefaultAcceptorFactory {
exec: Executor,
}
impl Service<SocketAddress> for DefaultAcceptorFactory {
type Output = TcpListener;
type Error = BoxError;
async fn serve(&self, addr: SocketAddress) -> Result<Self::Output, Self::Error> {
let acceptor = TcpListener::bind_address(addr, self.exec.clone()).await?;
Ok(acceptor)
}
}
pub trait Acceptor: Send + Sync + 'static {
type Stream: Io;
fn local_addr(&self) -> io::Result<SocketAddress>;
fn accept(self) -> impl Future<Output = Result<(Self::Stream, SocketAddress), Error>> + Send;
}
impl Acceptor for TcpListener {
type Stream = TcpStream;
fn local_addr(&self) -> io::Result<SocketAddress> {
Self::local_addr(self).map(Into::into)
}
#[inline]
async fn accept(self) -> Result<(Self::Stream, SocketAddress), Error> {
let (stream, addr) = Self::accept(&self).await.map_err(Error::io)?;
tracing::trace!(
network.peer.port = %addr.port,
network.peer.address = %addr.ip_addr,
"accepted incoming TCP connection"
);
Ok((stream, addr))
}
}
impl DefaultBinder {
#[must_use]
pub fn default_with_exec(exec: Executor) -> Self {
Self::new(
DefaultTimeout::new(DefaultAcceptorFactory::default(), Duration::from_secs(30)),
IoForwardService::new(exec),
)
}
}
impl Default for DefaultBinder {
fn default() -> Self {
Self::default_with_exec(Executor::default())
}
}
impl<S, F, StreamService> Socks5BinderSeal<S> for Binder<F, StreamService>
where
S: Io + Unpin,
F: SocketService<Socket: Acceptor<Stream: Unpin>>,
StreamService: Service<BridgeIo<S, <F::Socket as Acceptor>::Stream>, Error: Into<BoxError>>,
{
async fn accept_bind(
&self,
mut ingress_stream: S,
requested_bind_address: HostWithPort,
) -> Result<(), Error> {
tracing::trace!("socks5 server: bind: try to create acceptor @ {requested_bind_address}");
let HostWithPort {
host: requested_host,
port: requested_port,
} = requested_bind_address;
let Ok(requested_addr) = requested_host.try_as_ip() else {
tracing::debug!(
"bind command does not accept non-IP host {requested_host} as bind address"
);
let reply_kind = ReplyKind::AddressTypeNotSupported;
Reply::error_reply(reply_kind)
.write_to(&mut ingress_stream)
.await
.map_err(|err| Error::io(err).with_context("write server reply: bind failed"))?;
return Err(Error::aborted("bind failed").with_context(reply_kind));
};
let requested_address = SocketAddress::new(requested_addr, requested_port);
let bind_address = if let Some(bind_address) = self.bind_address {
tracing::trace!(
"socks5 server: bind: use server-defined bind interface: {bind_address}"
);
bind_address
} else {
tracing::debug!(
"socks5 server: bind: no server-defined bind interface: use requested client interface @ {requested_address}"
);
requested_address
};
let acceptor = match self.acceptor.bind_socket_with_address(bind_address).await {
Ok(twin) => twin,
Err(err) => {
let err = err.into();
tracing::debug!("make bind listener failed: {err:?}");
let reply_kind = ReplyKind::GeneralServerFailure;
Reply::error_reply(reply_kind)
.write_to(&mut ingress_stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: make bind listener failed")
})?;
return Err(Error::aborted("make bind listener failed")
.with_context(reply_kind)
.with_source(err));
}
};
let bind_address = match acceptor.local_addr() {
Ok(addr) => addr,
Err(err) => {
tracing::debug!(
"retrieve local addr of (tcp) acceptor failed @ {bind_address}: {err:?}",
);
let reply_kind = ReplyKind::GeneralServerFailure;
Reply::error_reply(reply_kind)
.write_to(&mut ingress_stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: make bind listener failed")
})?;
return Err(Error::aborted("make bind listener failed").with_context(reply_kind));
}
};
Reply::new(bind_address)
.write_to(&mut ingress_stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: bind: acceptor listener ready")
})?;
let accept_future = acceptor.accept();
let result = match self.accept_timeout {
Some(duration) => match tokio::time::timeout(duration, accept_future).await {
Ok(result) => result,
Err(err) => {
tracing::debug!("accept future timed out @ {bind_address}: {err:?}",);
let reply_kind = ReplyKind::TtlExpired;
Reply::error_reply(reply_kind)
.write_to(&mut ingress_stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: bind failed")
})?;
return Err(Error::aborted("bind failed").with_context(reply_kind));
}
},
None => accept_future.await,
};
let (incoming_stream, incoming_addr) = match result {
Ok((stream, addr)) => (stream, addr),
Err(err) => {
let err: BoxError = err.into();
tracing::debug!("socks5 server: abort: bind failed @ {bind_address}: {err:?}",);
let reply_kind = (&err).into();
Reply::error_reply(reply_kind)
.write_to(&mut ingress_stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: bind failed")
})?;
return Err(Error::aborted("bind failed")
.with_context(reply_kind)
.with_source(err));
}
};
tracing::trace!(
"incoming connection {incoming_addr} received on bind interface {bind_address}",
);
Reply::new(incoming_addr)
.write_to(&mut ingress_stream)
.await
.map_err(|err| {
Error::io(err).with_context("write server reply: bind: connection received")
})?;
tracing::trace!(
"socks5 server @ {bind_address}: bind: ready to serve from {incoming_addr}",
);
self.service
.serve(BridgeIo(ingress_stream, incoming_stream))
.instrument(tracing::trace_span!("socks5::bind::serve"))
.await
.map(drop)
.map_err(|err| Error::service(err).with_context("serve bind pipe"))
}
}
#[cfg(test)]
pub(crate) use test::MockBinder;
#[cfg(test)]
mod test {
#![expect(
clippy::unreachable,
reason = "test fixtures: arms gated on the mock variants the test sets up"
)]
use super::*;
use rama_net::address::HostWithPort;
use std::{ops::DerefMut, sync::Arc};
use tokio::sync::Mutex;
#[derive(Debug)]
pub(crate) struct MockBinder {
reply: MockReply,
}
#[derive(Debug)]
enum MockReply {
Success {
bind_addr: HostWithPort,
second_reply: MockSecondReply,
},
Error(ReplyKind),
}
#[derive(Debug)]
enum MockSecondReply {
Success {
recv_addr: HostWithPort,
target: Option<Arc<Mutex<tokio_test::io::Mock>>>,
},
Error(ReplyKind),
}
impl MockBinder {
pub(crate) fn new(bind_addr: HostWithPort, recv_addr: HostWithPort) -> Self {
Self {
reply: MockReply::Success {
bind_addr,
second_reply: MockSecondReply::Success {
recv_addr,
target: None,
},
},
}
}
pub(crate) fn new_err(reply: ReplyKind) -> Self {
Self {
reply: MockReply::Error(reply),
}
}
pub(crate) fn new_bind_err(bind_addr: HostWithPort, reply: ReplyKind) -> Self {
Self {
reply: MockReply::Success {
bind_addr,
second_reply: MockSecondReply::Error(reply),
},
}
}
pub(crate) fn with_proxy_data(mut self, target: tokio_test::io::Mock) -> Self {
self.reply = match self.reply {
MockReply::Success {
bind_addr,
second_reply:
MockSecondReply::Success {
recv_addr,
target: None,
},
} => MockReply::Success {
bind_addr,
second_reply: MockSecondReply::Success {
recv_addr,
target: Some(Arc::new(Mutex::new(target))),
},
},
MockReply::Error(_) | MockReply::Success { .. } => unreachable!(),
};
self
}
}
impl<S> Socks5BinderSeal<S> for MockBinder
where
S: Io + Unpin,
{
async fn accept_bind(
&self,
mut stream: S,
_requested_bind_address: HostWithPort,
) -> Result<(), Error> {
match &self.reply {
MockReply::Success {
bind_addr,
second_reply,
} => {
Reply::new(bind_addr.clone())
.write_to(&mut stream)
.await
.map_err(Error::io)?;
match second_reply {
MockSecondReply::Success { recv_addr, target } => {
Reply::new(recv_addr.clone())
.write_to(&mut stream)
.await
.map_err(Error::io)?;
if let Some(target) = target.as_ref() {
let mut target = target.lock().await;
match tokio::io::copy_bidirectional(&mut stream, target.deref_mut())
.await
{
Ok((bytes_copied_north, bytes_copied_south)) => {
tracing::trace!(
%bytes_copied_north,
%bytes_copied_south,
"(proxy) I/O stream forwarder finished"
);
Ok(())
}
Err(err) => {
if rama_net::conn::is_connection_error(&err) {
Ok(())
} else {
Err(Error::io(err))
}
}
}
} else {
Ok(())
}
}
MockSecondReply::Error(reply_kind) => {
Reply::error_reply(*reply_kind)
.write_to(&mut stream)
.await
.map_err(Error::io)?;
Err(Error::aborted("mock abort 2nd reply").with_context(*reply_kind))
}
}
}
MockReply::Error(reply_kind) => {
Reply::error_reply(*reply_kind)
.write_to(&mut stream)
.await
.map_err(Error::io)?;
Err(Error::aborted("mock abort 1st reply").with_context(*reply_kind))
}
}
}
}
}