use std::{
io::{self, Cursor},
net::SocketAddr,
sync::Mutex,
task::{Context, Poll, ready},
};
use bytes::{BufMut, BytesMut};
use log::{debug, trace};
use shadowsocks::{
context::SharedContext,
lookup_then,
net::{AddrFamily, ConnectOpts, TcpStream as ShadowTcpStream, UdpSocket as ShadowUdpSocket},
relay::{
socks5::{Address, UdpAssociateHeader},
udprelay::{DatagramReceive, DatagramSend, DatagramSocket},
},
};
use tokio::io::ReadBuf;
use super::{
OutboundProxyClient, OutboundProxyKind, TcpDialer, chain::connect_chain_for_udp_associate,
socks5::Socks5Negotiator,
};
struct Socks5UdpRelay {
#[allow(dead_code)]
assoc_tcp: ShadowTcpStream,
relay_addr: Address,
relay_socket_addr: Option<SocketAddr>,
}
impl Socks5UdpRelay {
fn relay_addr(&self) -> &Address {
&self.relay_addr
}
}
pub struct OutboundProxyDatagram {
socket: ShadowUdpSocket,
relays: Vec<Socks5UdpRelay>,
target: Address,
first_hop_addr: SocketAddr,
recv_scratch: Mutex<Vec<u8>>,
}
impl OutboundProxyDatagram {
pub async fn associate<D>(
client: &OutboundProxyClient,
context: &SharedContext,
dialer: &D,
connect_opts: &ConnectOpts,
target: Address,
) -> io::Result<Self>
where
D: TcpDialer + Sync,
{
let hops = client.hops();
if hops.is_empty() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"empty outbound proxy chain",
));
}
for hop in hops {
if !matches!(hop.kind, OutboundProxyKind::Socks5 { .. }) {
return Err(io::Error::new(
io::ErrorKind::Unsupported,
"outbound UDP relay requires every hop to be SOCKS5",
));
}
}
let socket = ShadowUdpSocket::connect_any_with_opts(AddrFamily::Ipv4, connect_opts).await?;
let local_udp_addr = socket.local_addr()?;
trace!("outbound udp local socket bound to {}", local_udp_addr);
let mut relays: Vec<Socks5UdpRelay> = Vec::with_capacity(hops.len());
let mut announce: Address = Address::from(local_udp_addr);
for (idx, hop) in hops.iter().enumerate() {
let mut tcp = connect_chain_for_udp_associate(hops, idx, dialer).await?;
let auth = match &hop.kind {
OutboundProxyKind::Socks5 { auth } => auth,
_ => unreachable!("guarded above"),
};
let resp = Socks5Negotiator::establish_udp_associate(&mut tcp, announce.clone(), auth)
.await
.map_err(io::Error::other)?;
let tcp = tcp.try_into_tcp().map_err(|_| {
io::Error::other(
"internal error: outbound UDP relay produced a non-TCP keep-alive stream",
)
})?;
let relay_addr = resp.address;
let relay_socket_addr = match relay_addr {
Address::SocketAddress(sa) => Some(sa),
Address::DomainNameAddress(ref name, port) => {
let (sa, _) = lookup_then!(context, name, port, |sa| {
Ok::<SocketAddr, io::Error>(sa)
})?;
Some(sa)
}
};
trace!(
"outbound udp hop {}: relay_addr={} (resolved={:?}), announce={}",
idx, relay_addr, relay_socket_addr, announce
);
announce = match relay_socket_addr {
Some(sa) => Address::from(sa),
None => relay_addr.clone(),
};
relays.push(Socks5UdpRelay {
assoc_tcp: tcp,
relay_addr,
relay_socket_addr,
});
}
let first_hop_addr = relays[0]
.relay_socket_addr
.ok_or_else(|| io::Error::other("first SOCKS5 UDP relay returned an unresolved domain"))?;
socket.connect(first_hop_addr).await?;
debug!(
"outbound udp chain established: {} hop(s), first_hop={}, target={}",
relays.len(),
first_hop_addr,
target,
);
Ok(Self {
socket,
relays,
target,
first_hop_addr,
recv_scratch: Mutex::new(Vec::new()),
})
}
pub fn hop_count(&self) -> usize {
self.relays.len()
}
fn encode(&self, payload: &[u8]) -> BytesMut {
let n = self.relays.len();
let mut total = payload.len();
for i in 0..n {
let addr = if i + 1 < n {
self.relays[i + 1].relay_addr()
} else {
&self.target
};
total += UdpAssociateHeader::new(0, addr.clone()).serialized_len();
}
let mut buf = BytesMut::with_capacity(total);
for i in 0..n {
let addr = if i + 1 < n {
self.relays[i + 1].relay_addr().clone()
} else {
self.target.clone()
};
UdpAssociateHeader::new(0, addr).write_to_buf(&mut buf);
}
buf.put_slice(payload);
buf
}
fn decode(&self, recv_buf: &[u8]) -> io::Result<(Address, usize)> {
let n = self.relays.len();
let mut cur = Cursor::new(recv_buf);
let mut last_addr: Option<Address> = None;
for _ in 0..n {
let fut = std::pin::pin!(UdpAssociateHeader::read_from(&mut cur));
match futures::FutureExt::now_or_never(fut) {
Some(Ok(h)) => last_addr = Some(h.address),
Some(Err(err)) => {
return Err(io::Error::other(format!(
"outbound udp: failed to parse SOCKS5 UDP header: {err}"
)));
}
None => unreachable!("Cursor read_from should be synchronous"),
}
}
let pos = cur.position() as usize;
Ok((last_addr.expect("relays is non-empty"), pos))
}
}
impl DatagramSocket for OutboundProxyDatagram {
fn local_addr(&self) -> io::Result<SocketAddr> {
self.socket.local_addr()
}
}
impl DatagramSend for OutboundProxyDatagram {
fn poll_send(&self, cx: &mut Context<'_>, buf: &[u8]) -> Poll<io::Result<usize>> {
let wire = self.encode(buf);
let sent = ready!(self.socket.poll_send(cx, &wire))?;
if sent == wire.len() {
Poll::Ready(Ok(buf.len()))
} else {
Poll::Ready(Err(io::Error::from(io::ErrorKind::WriteZero)))
}
}
fn poll_send_to(&self, cx: &mut Context<'_>, buf: &[u8], _target: SocketAddr) -> Poll<io::Result<usize>> {
let _ = _target;
self.poll_send(cx, buf)
}
fn poll_send_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
self.socket.poll_send_ready(cx)
}
}
impl DatagramReceive for OutboundProxyDatagram {
fn poll_recv(&self, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<()>> {
let mut scratch = self.recv_scratch.lock().expect("recv scratch poisoned");
let needed = buf.remaining() + 1024;
if scratch.len() < needed {
scratch.resize(needed, 0);
}
let mut tmp = ReadBuf::new(&mut scratch);
ready!(self.socket.poll_recv(cx, &mut tmp))?;
let n = tmp.filled().len();
let (_inner_src, pos) = self.decode(&scratch[..n])?;
let payload_len = n - pos;
if payload_len > buf.remaining() {
return Poll::Ready(Err(io::Error::from(io::ErrorKind::InvalidData)));
}
buf.put_slice(&scratch[pos..n]);
Poll::Ready(Ok(()))
}
fn poll_recv_from(&self, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll<io::Result<SocketAddr>> {
ready!(self.poll_recv(cx, buf))?;
Poll::Ready(Ok(self.first_hop_addr))
}
fn poll_recv_ready(&self, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
self.socket.poll_recv_ready(cx)
}
}