Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,6 @@ quinn = { version = "0.11", default-features = false, features = [
"runtime-tokio",
"rustls-ring",
] }
quinn-udp = "0.5.14"

# socks5
bytes = "1.11.0"
Expand All @@ -77,6 +76,15 @@ rpmalloc = { version = "0.2.2", optional = true }
jemallocator = { package = "tikv-jemallocator", version = "0.6.1", optional = true }
mimalloc = { version = "0.1.48", default-features = false, optional = true }

[target.'cfg(unix)'.dependencies]
libc = "0.2.189"

[target.'cfg(any(target_os = "linux", target_os = "android"))'.dependencies]
rustix = { version = "1.1.4", features = ["net"] }

[target.'cfg(not(any(target_os = "linux", target_os = "android")))'.dependencies]
quinn-udp = "0.5.14"

[target.'cfg(target_os = "linux")'.dependencies]
sysctl = "0.7.1"
rtnetlink = "0.18"
Expand Down
18 changes: 18 additions & 0 deletions src/connect.rs
Original file line number Diff line number Diff line change
Expand Up @@ -752,6 +752,24 @@ impl UdpConnector<'_> {
}
}

#[cfg(any(target_os = "linux", target_os = "android"))]
fn configure_udp_path(socket: &UdpSocket) -> std::io::Result<()> {
use rustix::net::sockopt::{
Ipv4PathMtuDiscovery, Ipv6PathMtuDiscovery, set_ip_mtu_discover, set_ipv6_mtu_discover,
};

// RFC 9298 forbids the proxy from introducing IP fragmentation. PROBE
// preserves datagram boundaries and reports an oversized send as EMSGSIZE.
// https://www.rfc-editor.org/rfc/rfc9298.html#section-6
if socket.peer_addr()?.is_ipv4() {
set_ip_mtu_discover(socket, Ipv4PathMtuDiscovery::PROBE)?;
} else {
set_ipv6_mtu_discover(socket, Ipv6PathMtuDiscovery::PROBE)?;
}
Ok(())
}

#[cfg(not(any(target_os = "linux", target_os = "android")))]
fn configure_udp_path(socket: &UdpSocket) -> std::io::Result<()> {
let state = quinn_udp::UdpSocketState::new(socket.into())?;
if state.may_fragment() {
Expand Down
21 changes: 21 additions & 0 deletions src/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,27 @@ use self::{
use crate::{AuthMode, BootArgs, Proxy, Result, connect::Connector};

const CONNECTION_DRAIN_TIMEOUT: Duration = Duration::from_secs(5);
// A bounded default that accommodates common QUIC packets without reserving a
// maximum-sized UDP datagram for every active relay.
const MAX_UDP_RELAY_PAYLOAD_SIZE: usize = 1_500;

fn is_oversized_datagram_error(error: &std_io::Error) -> bool {
let Some(code) = error.raw_os_error() else {
return false;
};
#[cfg(windows)]
{
code == 10040
}
#[cfg(unix)]
{
code == libc::EMSGSIZE
}
#[cfg(not(any(unix, windows)))]
{
false
}
}

/// Trait for connection acceptors that handle incoming TCP streams.
pub trait Acceptor {
Expand Down
37 changes: 13 additions & 24 deletions src/server/masque.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,15 +36,17 @@ use tokio::{
time::timeout,
};

use super::{AuthMode, Context, Handle, Server, http::genca};
use super::{
AuthMode, Context, Handle, MAX_UDP_RELAY_PAYLOAD_SIZE, Server, http::genca,
is_oversized_datagram_error,
};
use crate::{connect::Connector, ext::Extension};

const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
const REQUEST_TIMEOUT: Duration = Duration::from_secs(10);
const SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(5);
const MAX_UDP_PAYLOAD: usize = 1_500;
const DATAGRAM_CAPSULE_TYPE: u64 = 0;
const MAX_CAPSULE_DATAGRAM_SIZE: usize = MAX_UDP_PAYLOAD + 8;
const MAX_CAPSULE_DATAGRAM_SIZE: usize = MAX_UDP_RELAY_PAYLOAD_SIZE + 8;

type BoxError = Box<dyn std::error::Error + Send + Sync>;
type Sessions = Arc<RwLock<HashMap<StreamId, Arc<Session>>>>;
Expand Down Expand Up @@ -654,7 +656,7 @@ where
// RFC 9298 forbids the proxy from fragmenting UDP payloads. One extra byte
// lets recv detect and discard datagrams above the common Ethernet MTU.
// https://www.rfc-editor.org/rfc/rfc9298.html#section-6
let mut payload = [0; MAX_UDP_PAYLOAD + 1];
let mut payload = [0; MAX_UDP_RELAY_PAYLOAD_SIZE + 1];
let mut capsules = CapsuleDecoder::default();
loop {
tokio::select! {
Expand Down Expand Up @@ -742,7 +744,7 @@ async fn forward_capsule_datagram(socket: &UdpSocket, datagram: &[u8]) -> io::Re
"missing CONNECT-UDP context ID",
));
};
if context_id != 0 || payload.len() > MAX_UDP_PAYLOAD {
if context_id != 0 || payload.len() > MAX_UDP_RELAY_PAYLOAD_SIZE {
return Ok(());
}
match socket.send(payload).await {
Expand All @@ -758,19 +760,6 @@ fn decode_capsule_header(data: &[u8]) -> Option<(u64, u64)> {
Some((capsule_type, capsule_length))
}

fn is_oversized_datagram_error(error: &io::Error) -> bool {
error.raw_os_error().is_some_and(|code| {
#[cfg(windows)]
{
code == 10040
}
#[cfg(not(windows))]
{
code == 90
}
})
}

fn encode_target_datagram(stream_id: StreamId, payload: &[u8]) -> io::Result<Bytes> {
// RFC 9297 prefixes an HTTP Datagram with the request's Quarter Stream ID.
// https://www.rfc-editor.org/rfc/rfc9297.html#section-2.1
Expand All @@ -785,7 +774,7 @@ fn encode_target_datagram(stream_id: StreamId, payload: &[u8]) -> io::Result<Byt
}

fn usable_udp_payload(payload: &[u8], length: usize) -> Option<&[u8]> {
(length <= MAX_UDP_PAYLOAD).then(|| &payload[..length])
(length <= MAX_UDP_RELAY_PAYLOAD_SIZE).then(|| &payload[..length])
}

async fn process_datagrams<H>(
Expand Down Expand Up @@ -815,7 +804,7 @@ where
if context_id != 0 {
continue;
}
if payload.len() > MAX_UDP_PAYLOAD {
if payload.len() > MAX_UDP_RELAY_PAYLOAD_SIZE {
continue;
}
tracing::trace!(
Expand Down Expand Up @@ -1010,13 +999,13 @@ mod tests {

#[test]
fn drops_udp_payloads_above_the_proxy_limit() {
let payload = [0; MAX_UDP_PAYLOAD + 1];
let payload = [0; MAX_UDP_RELAY_PAYLOAD_SIZE + 1];
assert_eq!(
usable_udp_payload(&payload, MAX_UDP_PAYLOAD)
usable_udp_payload(&payload, MAX_UDP_RELAY_PAYLOAD_SIZE)
.expect("payload at the limit is accepted")
.len(),
MAX_UDP_PAYLOAD
MAX_UDP_RELAY_PAYLOAD_SIZE
);
assert!(usable_udp_payload(&payload, MAX_UDP_PAYLOAD + 1).is_none());
assert!(usable_udp_payload(&payload, MAX_UDP_RELAY_PAYLOAD_SIZE + 1).is_none());
}
}
143 changes: 127 additions & 16 deletions src/server/socks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ use std::{
},
};

use bytes::BytesMut;
use tokio::{
io::AsyncWriteExt,
net::{TcpListener, TcpStream, UdpSocket},
Expand All @@ -26,9 +27,12 @@ use self::{
connect::{self, Connect},
},
error::Error,
proto::{Address, Reply, UdpHeader},
proto::{Address, Reply},
};
use super::{
Acceptor, Context, Handle, MAX_UDP_RELAY_PAYLOAD_SIZE, Server, drain_connections, io,
is_oversized_datagram_error, log_connection_result,
};
use super::{Acceptor, Context, Handle, Server, drain_connections, io, log_connection_result};
use crate::connect::{Connector, TcpConnector, UdpConnector};

/// SOCKS5 acceptor.
Expand Down Expand Up @@ -212,16 +216,18 @@ async fn handle_connect(
}
}

const MAX_UDP_RELAY_PACKET_SIZE: usize = 1500;
// Keep the default bounded while allowing common 1,200-1,350 byte QUIC packets.
// AssociatedUdpSocket adds the largest RFC 1928 header and one truncation
// detection byte to its reusable receive buffer.
// https://www.rfc-editor.org/rfc/rfc1928.html#section-7
const UDP_PAYLOAD_RECV_BUFFER_SIZE: usize = MAX_UDP_RELAY_PAYLOAD_SIZE + 1;

#[instrument(skip(associate, connector), level = Level::DEBUG)]
async fn handle_udp(
associate: UdpAssociate<associate::NeedReply>,
address: Address,
connector: UdpConnector<'_>,
) -> std::io::Result<()> {
const BUF_SIZE: usize = MAX_UDP_RELAY_PACKET_SIZE - UdpHeader::max_serialized_len();

let socket = UdpSocket::bind(SocketAddr::from((associate.local_addr()?.ip(), 0))).await?;
let listen_addr = socket.local_addr()?;
tracing::info!("[SOCKS5][UDP] listening on: {listen_addr}");
Expand All @@ -230,7 +236,12 @@ async fn handle_udp(
.reply(Reply::Succeeded, Address::from(listen_addr))
.await?;

let inbound = AssociatedUdpSocket::from((socket, BUF_SIZE));
let inbound = AssociatedUdpSocket::new(socket, MAX_UDP_RELAY_PAYLOAD_SIZE)?;
let mut inbound_buffer = vec![0; inbound.recv_buffer_size()];
let mut preferred_buffer = [0; UDP_PAYLOAD_RECV_BUFFER_SIZE];
let mut fallback_buffer = [0; UDP_PAYLOAD_RECV_BUFFER_SIZE];
let mut preferred_response = BytesMut::new();
let mut fallback_response = BytesMut::new();
let (preferred_outbound, fallback_outbound) = connector.create_socket_dual_stack().await?;

// Determine the source IP for UDP packets:
Expand All @@ -248,8 +259,7 @@ async fn handle_udp(
loop {
let result = tokio::select! {
req = async {
inbound.set_max_packet_size(BUF_SIZE);
let (pkt, frag, dst_addr, src_addr) = inbound.recv_from().await?;
let (pkt, frag, dst_addr, src_addr) = inbound.recv_from(&mut inbound_buffer).await?;

if frag != 0 {
return Err(Error::from("[SOCKS5][UDP] packet fragment is not supported"));
Expand Down Expand Up @@ -293,13 +303,13 @@ async fn handle_udp(
Address::SocketAddress(target_addr) => {
tracing::info!("[SOCKS5][UDP] {src_addr} -> {target_addr} forwarding packet, size {}", pkt.len());
connector
.send_packet(&pkt, target_addr, &preferred_outbound, fallback_outbound.as_ref())
.send_packet(pkt, target_addr, &preferred_outbound, fallback_outbound.as_ref())
.await?;
}
Address::DomainAddress(domain, port) => {
tracing::info!("[SOCKS5][UDP] {src_addr} -> {domain}:{port} forwarding packet, size {}", pkt.len());
connector
.send_packet(&pkt, (domain, port), &preferred_outbound, fallback_outbound.as_ref())
.send_packet(pkt, (domain, port), &preferred_outbound, fallback_outbound.as_ref())
.await?;
}
}
Expand All @@ -308,29 +318,55 @@ async fn handle_udp(
} => req,

preferred_resp = async {
let mut buf = [0u8; MAX_UDP_RELAY_PACKET_SIZE];
let (len, remote_addr) = preferred_outbound.recv_from(&mut buf).await?;
let (len, remote_addr) = match preferred_outbound.recv_from(&mut preferred_buffer).await {
Ok(received) => received,
Err(error) if is_oversized_datagram_error(&error) => return Ok(()),
Err(error) => return Err(Error::from(error)),
};
if len > MAX_UDP_RELAY_PAYLOAD_SIZE {
tracing::trace!("[SOCKS5][UDP] dropping oversized packet from {remote_addr}");
return Ok(());
}
let src_addr = SocketAddr::new(src_ip, src_port.load(Ordering::Relaxed));

tracing::info!("[SOCKS5][UDP] {src_addr} <- {remote_addr} feedback to incoming, packet size {len}");

inbound
.send_to(&buf[..len], 0, remote_addr.into(), src_addr)
.send_to_buffered(
&preferred_buffer[..len],
0,
remote_addr.into(),
src_addr,
&mut preferred_response,
)
.await
.map(|_| ())
.map_err(Error::from)
} => preferred_resp,

fallback_resp = async {
if let Some(ref fallback_outbound) = fallback_outbound {
let mut buf = [0u8; MAX_UDP_RELAY_PACKET_SIZE];
let (len, remote_addr) = fallback_outbound.recv_from(&mut buf).await?;
let (len, remote_addr) = match fallback_outbound.recv_from(&mut fallback_buffer).await {
Ok(received) => received,
Err(error) if is_oversized_datagram_error(&error) => return Ok(()),
Err(error) => return Err(Error::from(error)),
};
if len > MAX_UDP_RELAY_PAYLOAD_SIZE {
tracing::trace!("[SOCKS5][UDP] dropping oversized packet from {remote_addr}");
return Ok(());
}
let src_addr = SocketAddr::new(src_ip, src_port.load(Ordering::Relaxed));

tracing::info!("[SOCKS5][UDP] {src_addr} <- {remote_addr} feedback to incoming, packet size {len}");

inbound
.send_to(&buf[..len], 0, remote_addr.into(), src_addr)
.send_to_buffered(
&fallback_buffer[..len],
0,
remote_addr.into(),
src_addr,
&mut fallback_response,
)
.await
.map(|_| ())
.map_err(Error::from)
Expand Down Expand Up @@ -454,3 +490,78 @@ async fn handle_bind(
}
}
}

#[cfg(test)]
mod tests {
use std::{net::SocketAddr, time::Duration};

use tokio::{net::UdpSocket, time::timeout};

use super::{Address, AssociatedUdpSocket, MAX_UDP_RELAY_PAYLOAD_SIZE};

#[tokio::test]
async fn udp_relay_buffer_accepts_a_quic_sized_datagram() {
const QUIC_PACKET_SIZE: usize = 1350;

let receiver_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let receiver_address = receiver_socket.local_addr().unwrap();
let receiver =
AssociatedUdpSocket::new(receiver_socket, MAX_UDP_RELAY_PAYLOAD_SIZE).unwrap();
let mut receive_buffer = vec![0; receiver.recv_buffer_size()];

let sender_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let sender = AssociatedUdpSocket::new(sender_socket, MAX_UDP_RELAY_PAYLOAD_SIZE).unwrap();
let destination = Address::from(SocketAddr::from(([127, 0, 0, 1], 443)));
let payload = vec![0x5a; QUIC_PACKET_SIZE];

let sent = sender
.send_to(&payload, 0, destination.clone(), receiver_address)
.await
.unwrap();
assert_eq!(sent, payload.len());

let (received, fragment, received_destination, _) = timeout(
Duration::from_secs(1),
receiver.recv_from(&mut receive_buffer),
)
.await
.unwrap()
.unwrap();
assert_eq!(fragment, 0);
assert_eq!(received_destination, destination);
assert_eq!(received, payload);
}

#[tokio::test]
async fn udp_relay_drops_an_oversized_datagram_without_forwarding_a_truncated_packet() {
let receiver_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let receiver_address = receiver_socket.local_addr().unwrap();
let receiver =
AssociatedUdpSocket::new(receiver_socket, MAX_UDP_RELAY_PAYLOAD_SIZE).unwrap();
let mut receive_buffer = vec![0; receiver.recv_buffer_size()];

let sender_socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let sender = AssociatedUdpSocket::new(sender_socket, receiver.recv_buffer_size()).unwrap();
let destination = Address::from(SocketAddr::from(([127, 0, 0, 1], 443)));
let oversized = vec![0x5a; receiver.recv_buffer_size()];
sender
.send_to(&oversized, 0, destination.clone(), receiver_address)
.await
.unwrap();
sender
.send_to(b"valid", 0, destination.clone(), receiver_address)
.await
.unwrap();

let (received, fragment, received_destination, _) = timeout(
Duration::from_secs(1),
receiver.recv_from(&mut receive_buffer),
)
.await
.unwrap()
.unwrap();
assert_eq!(fragment, 0);
assert_eq!(received_destination, destination);
assert_eq!(received, b"valid");
}
}
Loading
Loading