use std::{
future::Future,
io,
net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr},
sync::Arc,
time::Duration,
};
use bytes::{Buf, BytesMut};
use futures::{Sink, SinkExt, Stream, StreamExt};
use tokio::{
io::{AsyncRead, AsyncWrite},
net::UdpSocket,
sync::mpsc::UnboundedSender,
task::JoinHandle,
};
use tokio_util::{codec::LengthDelimitedCodec, sync::CancellationToken};
use tracing::*;
use crate::connection::BridgeConn;
use crate::connection::make_socket;
use crate::error::TransportError;
const ETHERNET_V2_MTU: u16 = 1500;
const LENGTH_DELIMITER_BYTELEN: usize = 2;
const INITIAL_CONNECTION_TIMEOUT: Duration = Duration::from_secs(10);
trait UdpPeer: Send + Sync + 'static {
async fn recv(
&self,
sock: &UdpSocket,
buf: &mut BytesMut,
fwd_addr: SocketAddr,
) -> io::Result<(usize, SocketAddr)>;
async fn send(&self, sock: &UdpSocket, data: &[u8], dest: SocketAddr) -> io::Result<usize>;
}
struct ConnectedPeer;
impl UdpPeer for ConnectedPeer {
async fn recv(
&self,
sock: &UdpSocket,
buf: &mut BytesMut,
fwd_addr: SocketAddr,
) -> io::Result<(usize, SocketAddr)> {
let len = sock.recv_buf(buf).await?;
Ok((len, fwd_addr))
}
async fn send(&self, sock: &UdpSocket, data: &[u8], _dest: SocketAddr) -> io::Result<usize> {
sock.send(data).await
}
}
struct SharedPeer;
impl UdpPeer for SharedPeer {
async fn recv(
&self,
sock: &UdpSocket,
buf: &mut BytesMut,
_fwd_addr: SocketAddr,
) -> io::Result<(usize, SocketAddr)> {
sock.recv_buf_from(buf).await
}
async fn send(&self, sock: &UdpSocket, data: &[u8], dest: SocketAddr) -> io::Result<usize> {
sock.send_to(data, dest).await
}
}
fn address_match(original: SocketAddr, incoming: SocketAddr) -> bool {
if incoming == original {
true
} else {
match (original.ip(), incoming.ip()) {
(IpAddr::V4(orig), IpAddr::V6(_)) => {
SocketAddr::from((orig.to_ipv6_mapped(), original.port())) == incoming
}
(IpAddr::V6(_), IpAddr::V4(inc)) => {
original == SocketAddr::from((inc.to_ipv6_mapped(), incoming.port()))
}
_ => false,
}
}
}
async fn udp_to_transport_task<W, P, TL>(
peer: P,
sock: Arc<UdpSocket>,
mut framed_writer: W,
fwd_addr: SocketAddr,
tr_label: TL,
mtu: u16,
token: CancellationToken,
) -> Result<(), io::Error>
where
W: Sink<bytes::Bytes, Error = io::Error> + Unpin + Send,
P: UdpPeer,
TL: std::fmt::Display,
{
let mut dn_buf = BytesMut::with_capacity(mtu as usize);
loop {
tokio::select! {
res = peer.recv(&sock, &mut dn_buf, fwd_addr) => {
let (len, src) = res.map_err(|e| {
error!("error receiving from forward socket: {e}");
e
})?;
if !address_match(fwd_addr, src) {
debug!("received {len}B from alt addr {src} -- ignoring");
dn_buf.clear();
if !dn_buf.try_reclaim(mtu as usize) {
warn!("unable to reclaim bytes in buffer: {} ", dn_buf.capacity());
}
continue;
}
trace!(" <-{fwd_addr} read {len}B");
framed_writer.send(dn_buf.copy_to_bytes(len)).await.map_err(|e| {
error!("error sending to transport connection: {e}");
e
})?;
trace!(" {tr_label}<- wrote {len}B");
dn_buf.clear();
if !dn_buf.try_reclaim(mtu as usize) {
warn!("unable to reclaim bytes in buffer: {} ", dn_buf.capacity());
}
}
_ = token.cancelled() => {
debug!("end io copy from {fwd_addr}<->{tr_label}");
break;
}
}
}
Ok(())
}
async fn transport_to_udp_task<R, P, TL>(
peer: P,
mut framed_reader: R,
sock: Arc<UdpSocket>,
fwd_addr: SocketAddr,
tr_label: TL,
token: CancellationToken,
) -> Result<(), io::Error>
where
R: Stream<Item = Result<bytes::BytesMut, io::Error>> + Unpin + Send,
P: UdpPeer,
TL: std::fmt::Display,
{
loop {
tokio::select! {
res = framed_reader.next() => {
match res {
None => {
info!("connection closed");
break;
}
Some(Ok(buf)) => {
let len = buf.len();
trace!("{tr_label}-> read {len}B");
let mut sent = 0;
let mut sends = 1;
while sent < len {
let len_sent = peer.send(&sock, &buf[sent..len], fwd_addr).await.map_err(|e| {
error!("error sending to egress socket: {e}");
e
})?;
sent += len_sent;
trace!(" ->{fwd_addr} wrote {len_sent}B {sends} send");
sends +=1;
}
}
Some(Err(e)) => {
error!("error reading from transport conn: {e}");
return Err(e);
}
}
}
_ = token.cancelled() => {
debug!("end io copy from {fwd_addr}<->{tr_label}");
break;
}
}
}
Ok(())
}
async fn run_forward_pair<F1, F2>(recv_task: F1, send_task: F2, token: CancellationToken)
where
F1: Future<Output = Result<(), io::Error>> + Send + 'static,
F2: Future<Output = Result<(), io::Error>> + Send + 'static,
{
let mut tasks = tokio::task::JoinSet::new();
tasks.spawn(recv_task);
tasks.spawn(send_task);
let mut token = Some(token);
while let Some(res) = tasks.join_next().await {
if let Err(err) = res {
error!("bridge udp forwarder join error: {err}");
} else if let Ok(Err(err)) = res {
error!("bridge udp forwarder error: {err}");
}
if let Some(token) = token.take() {
token.cancel();
}
}
}
pub struct UdpForwarder {}
impl UdpForwarder {
pub async fn launch_initiator(
egress_conn: BridgeConn,
bind_addr: Option<SocketAddr>,
close_tx: Option<UnboundedSender<()>>,
token: CancellationToken,
) -> Result<(SocketAddr, JoinHandle<()>), TransportError> {
let bind_addr = bind_addr.unwrap_or(match egress_conn.endpoint.is_ipv4() {
true => (Ipv4Addr::LOCALHOST, 0).into(),
false => (Ipv6Addr::LOCALHOST, 0).into(),
});
let socket = make_socket(Some(bind_addr)).map_err(TransportError::SocketIo)?;
let socket = Arc::new(UdpSocket::from_std(socket).map_err(TransportError::SocketIo)?);
let local_addr = socket.local_addr().map_err(TransportError::SocketIo)?;
info!("udp forwarder started listening on: {local_addr}",);
Ok((
local_addr,
tokio::spawn(initiator::process_udp(
egress_conn.reader,
egress_conn.writer,
socket.clone(),
ETHERNET_V2_MTU,
close_tx,
token,
)),
))
}
}
pub mod initiator {
use super::*;
pub async fn process_udp<R, W>(
reader: R,
writer: W,
sock: Arc<UdpSocket>,
mtu: u16,
close_tx: Option<UnboundedSender<()>>,
token: CancellationToken,
) where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
info!("starting udp forward");
let mut dn_buf = BytesMut::with_capacity(mtu as usize);
let mut framed_writer = LengthDelimitedCodec::builder()
.length_field_length(LENGTH_DELIMITER_BYTELEN)
.new_write(writer);
let framed_reader = LengthDelimitedCodec::builder()
.length_field_length(LENGTH_DELIMITER_BYTELEN)
.new_read(reader);
let fwd_initial_recv_fut =
tokio::time::timeout(INITIAL_CONNECTION_TIMEOUT, sock.recv_buf_from(&mut dn_buf));
let fwd_addr = match token.run_until_cancelled(fwd_initial_recv_fut).await {
Some(res) => {
match res {
Ok(Ok((len, src))) => {
trace!(" <- [fw] read {len}B");
if let Err(e) = framed_writer.send(dn_buf.copy_to_bytes(len)).await {
debug!("error sending to transport connection: {e}");
None
} else {
trace!("[tr] <- wrote {len}B");
Some(src)
}
}
Ok(Err(e)) => {
debug!("error receiving from egress socket: {e}");
None
}
Err(_) => {
debug!("forwarder timed out");
None
}
}
}
None => {
debug!("forwarder cancelled before initial receive");
None
}
};
let Some(fwd_addr) = fwd_addr else {
if let Some(tx) = close_tx {
tx.send(()).ok();
}
return;
};
if let Err(e) = sock.connect(fwd_addr).await {
error!("udp sock config failure: {e}");
if let Some(tx) = close_tx {
tx.send(()).ok();
}
return;
}
run_forward_pair(
udp_to_transport_task(
ConnectedPeer,
sock.clone(),
framed_writer,
fwd_addr,
"[tr]",
mtu,
token.clone(),
),
transport_to_udp_task(
ConnectedPeer,
framed_reader,
sock.clone(),
fwd_addr,
"[tr]",
token.clone(),
),
token,
)
.await;
if let Some(tx) = close_tx {
tx.send(()).ok();
}
info!("transport udp forwarder shutdown");
}
}
pub mod responder {
use super::*;
use crate::session::Session;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::ReadBuf;
pub async fn process_udp<R, W>(
rd: R,
wr: W,
sock: Arc<UdpSocket>,
session: Session,
mtu: u16,
token: CancellationToken,
) where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
let tr_addr = session.transport_remote();
let fw_addr = session.forward_remote();
let local_fw_addr = sock.local_addr().unwrap();
info!("starting udp forward {tr_addr:?}->([tr_local] -> {local_fw_addr:?}) -> {fw_addr:?}");
let framed_writer = LengthDelimitedCodec::builder()
.length_field_length(LENGTH_DELIMITER_BYTELEN)
.new_write(LoggingIo::new(wr, "".into()));
let framed_reader = LengthDelimitedCodec::builder()
.length_field_length(LENGTH_DELIMITER_BYTELEN)
.new_read(LoggingIo::new(rd, "".into()));
run_forward_pair(
udp_to_transport_task(
SharedPeer,
sock.clone(),
framed_writer,
fw_addr,
tr_addr,
mtu,
token.clone(),
),
transport_to_udp_task(
SharedPeer,
framed_reader,
sock.clone(),
fw_addr,
tr_addr,
token.clone(),
),
token,
)
.await;
drop(sock);
}
struct LoggingIo<T> {
inner: T,
name: String,
}
impl<T> LoggingIo<T> {
fn new(inner: T, name: String) -> Self {
Self { inner, name }
}
}
impl<T: AsyncRead + Unpin> AsyncRead for LoggingIo<T> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<tokio::io::Result<()>> {
let before = buf.filled().len();
let result = Pin::new(&mut self.inner).poll_read(cx, buf);
if let Poll::Ready(Ok(())) = &result {
let bytes_read = buf.filled().len() - before;
if bytes_read > 0 {
trace!("{}: Read {} bytes", self.name, bytes_read,);
}
}
result
}
}
impl<T: AsyncWrite + Unpin> AsyncWrite for LoggingIo<T> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, tokio::io::Error>> {
let result = Pin::new(&mut self.inner).poll_write(cx, buf);
if let Poll::Ready(Ok(bytes_written)) = &result
&& *bytes_written > 0
{
trace!("{}: Wrote {} bytes", self.name, bytes_written,);
}
result
}
fn poll_flush(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), tokio::io::Error>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), tokio::io::Error>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
#[cfg(test)]
mod test {
use super::*;
use std::net::SocketAddr;
#[test]
fn addr_equality() {
let addr1: SocketAddr = "[::ffff:178.79.168.250]:51822".parse().unwrap();
let addr2: SocketAddr = "178.79.168.250:51822".parse().unwrap();
assert_ne!(addr1, addr2);
assert!(address_match(addr1, addr2));
assert!(address_match(addr2, addr1));
assert!(address_match(addr1, addr1));
assert!(address_match(addr2, addr2));
let addr3: SocketAddr = "192.168.1.1:51822".parse().unwrap(); let addr4: SocketAddr = "178.79.168.250:9000".parse().unwrap();
assert!(!address_match(addr1, addr3));
assert!(!address_match(addr3, addr1));
assert!(!address_match(addr2, addr3));
assert!(!address_match(addr3, addr2));
assert!(!address_match(addr1, addr4));
assert!(!address_match(addr4, addr1));
assert!(!address_match(addr2, addr4));
assert!(!address_match(addr4, addr2));
}
}
}