Redesign runtime w/ include-aware config + Module Split + Listener Lifecycle + Atomic Reload

This commit is contained in:
Alexey
2026-08-22 16:13:25 +03:00
parent f9910fb29e
commit 7f4b87bea4
72 changed files with 14433 additions and 12976 deletions
+122 -34
View File
@@ -4,6 +4,83 @@ use std::os::fd::{AsRawFd, BorrowedFd, RawFd};
use tokio::io::Interest;
use tokio::io::unix::AsyncFd;
struct TcpMaxSegmentGuard {
fd: RawFd,
original: libc::c_int,
changed: bool,
}
fn tcp_max_segment(fd: RawFd) -> Result<libc::c_int> {
let mut value: libc::c_int = 0;
let mut length = std::mem::size_of::<libc::c_int>() as libc::socklen_t;
let rc = unsafe {
libc::getsockopt(
fd,
libc::IPPROTO_TCP,
libc::TCP_MAXSEG,
&mut value as *mut libc::c_int as *mut libc::c_void,
&mut length,
)
};
if rc != 0 {
return Err(Error::last_os_error());
}
Ok(value)
}
fn set_tcp_max_segment(fd: RawFd, value: libc::c_int) -> Result<()> {
let rc = unsafe {
libc::setsockopt(
fd,
libc::IPPROTO_TCP,
libc::TCP_MAXSEG,
&value as *const libc::c_int as *const libc::c_void,
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
if rc != 0 {
return Err(Error::last_os_error());
}
Ok(())
}
impl TcpMaxSegmentGuard {
fn install(fd: RawFd, requested: usize) -> Result<Self> {
let requested = libc::c_int::try_from(requested).map_err(|_| {
Error::new(
ErrorKind::InvalidInput,
"TCP fragment size exceeds the platform integer range",
)
})?;
let original = tcp_max_segment(fd)?;
let changed = requested < original;
if changed {
set_tcp_max_segment(fd, requested)?;
}
Ok(Self {
fd,
original,
changed,
})
}
fn restore(mut self) -> Result<()> {
if self.changed {
set_tcp_max_segment(self.fd, self.original)?;
self.changed = false;
}
Ok(())
}
}
impl Drop for TcpMaxSegmentGuard {
fn drop(&mut self) {
if self.changed {
let _ = set_tcp_max_segment(self.fd, self.original);
}
}
}
fn force_tcp_push(fd: RawFd) -> Result<()> {
let enabled: libc::c_int = 1;
let rc = unsafe {
@@ -25,8 +102,9 @@ fn force_tcp_push(fd: RawFd) -> Result<()> {
///
/// `fd` must refer to a connected, nonblocking TCP socket and remain valid for
/// the duration of this call. The caller retains ownership of the original fd.
/// `MSG_EOR` is only a best-effort Linux hint for TCP: offloads, loss, and
/// retransmission may coalesce these write boundaries on the wire.
/// The accepted socket is clamped to the requested `TCP_MAXSEG` for this send
/// and restored on success, error, or task cancellation. `MSG_EOR` remains only
/// a best-effort Linux hint: offloads may still coalesce capture boundaries.
pub(crate) async fn send_tcp_fragmented_fd(
fd: RawFd,
data: &[u8],
@@ -49,42 +127,52 @@ pub(crate) async fn send_tcp_fragmented_fd(
let borrowed_fd = unsafe { BorrowedFd::borrow_raw(fd) };
let duplicated_fd = borrowed_fd.try_clone_to_owned()?;
let async_fd = AsyncFd::with_interest(duplicated_fd, Interest::WRITABLE)?;
let mss_guard = TcpMaxSegmentGuard::install(async_fd.get_ref().as_raw_fd(), fragment_size)?;
for fragment in data.chunks(fragment_size) {
let mut offset = 0;
while offset < fragment.len() {
let mut writable = async_fd.writable().await?;
let sent = match writable.try_io(|inner| {
let remaining = &fragment[offset..];
let sent = unsafe {
libc::send(
inner.get_ref().as_raw_fd(),
remaining.as_ptr().cast::<libc::c_void>(),
remaining.len(),
libc::MSG_DONTWAIT | libc::MSG_EOR | libc::MSG_NOSIGNAL,
)
let send_result = async {
for fragment in data.chunks(fragment_size) {
let mut offset = 0;
while offset < fragment.len() {
let mut writable = async_fd.writable().await?;
let sent = match writable.try_io(|inner| {
let remaining = &fragment[offset..];
let sent = unsafe {
libc::send(
inner.get_ref().as_raw_fd(),
remaining.as_ptr().cast::<libc::c_void>(),
remaining.len(),
libc::MSG_DONTWAIT | libc::MSG_EOR | libc::MSG_NOSIGNAL,
)
};
if sent < 0 {
Err(Error::last_os_error())
} else if sent == 0 {
Err(Error::new(
ErrorKind::WriteZero,
"fragmented TCP send returned zero",
))
} else {
Ok(sent as usize)
}
}) {
Ok(Ok(sent)) => sent,
Ok(Err(error)) if error.kind() == ErrorKind::Interrupted => continue,
Ok(Err(error)) => return Err(error),
Err(_) => continue,
};
if sent < 0 {
Err(Error::last_os_error())
} else if sent == 0 {
Err(Error::new(
ErrorKind::WriteZero,
"fragmented TCP send returned zero",
))
} else {
Ok(sent as usize)
}
}) {
Ok(Ok(sent)) => sent,
Ok(Err(error)) if error.kind() == ErrorKind::Interrupted => continue,
Ok(Err(error)) => return Err(error),
Err(_) => continue,
};
offset += sent;
force_tcp_push(async_fd.get_ref().as_raw_fd())?;
offset += sent;
force_tcp_push(async_fd.get_ref().as_raw_fd())?;
}
}
Ok(())
}
.await;
Ok(())
let restore_result = mss_guard.restore();
send_result.and(restore_result)
}
#[cfg(test)]
#[path = "fragmented_send_wire_tests.rs"]
mod wire_tests;
@@ -0,0 +1,232 @@
use std::net::SocketAddr;
use std::os::fd::{AsRawFd, BorrowedFd};
use std::path::Path;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use crate::config::ProxyConfig;
use crate::crypto::{SecureRandom, sha256_hmac};
use crate::error::HandshakeResult;
use crate::protocol::constants::TLS_VERSION;
use crate::protocol::tls;
use crate::proxy::handshake::{
TlsResponseWriteOptions, handle_tls_handshake_with_shared_and_options,
};
use crate::proxy::shared_state::ProxySharedState;
use crate::stats::ReplayChecker;
use crate::transport::socket::{ListenOptions, create_listener};
const SECRET: [u8; 16] = [0x56; 16];
const SECRET_HEX: &str = "56565656565656565656565656565656";
const BULK_PAYLOAD_LEN: usize = 8192;
fn make_valid_tls_client_hello(tls_len: usize) -> Vec<u8> {
const TLS_AES_128_GCM_SHA256: [u8; 2] = [0x13, 0x01];
const TLS_EXTENSION_KEY_SHARE: u16 = 0x0033;
const TLS_EXTENSION_PADDING: u16 = 0x0015;
const X25519_KEY_SHARE_LEN: usize = 32;
let fill = 0x42_u8;
let session_id_len = 32_usize;
let mut extensions = Vec::new();
let mut key_share = Vec::new();
key_share.extend_from_slice(&tls::TLS_NAMED_GROUP_X25519.to_be_bytes());
key_share.extend_from_slice(&(X25519_KEY_SHARE_LEN as u16).to_be_bytes());
key_share.push(9);
key_share.resize(key_share.len() + X25519_KEY_SHARE_LEN - 1, 0);
let mut key_share_extension = Vec::new();
key_share_extension.extend_from_slice(&(key_share.len() as u16).to_be_bytes());
key_share_extension.extend_from_slice(&key_share);
extensions.extend_from_slice(&TLS_EXTENSION_KEY_SHARE.to_be_bytes());
extensions.extend_from_slice(&(key_share_extension.len() as u16).to_be_bytes());
extensions.extend_from_slice(&key_share_extension);
let base_tls_len = 4
+ 2
+ 32
+ 1
+ session_id_len
+ 2
+ TLS_AES_128_GCM_SHA256.len()
+ 1
+ 1
+ 2
+ extensions.len();
let padding_len = tls_len
.checked_sub(base_tls_len + 4)
.expect("wire ClientHello must leave room for padding");
extensions.extend_from_slice(&TLS_EXTENSION_PADDING.to_be_bytes());
extensions.extend_from_slice(&(padding_len as u16).to_be_bytes());
extensions.resize(extensions.len() + padding_len, fill);
let body_len = tls_len - 4;
let mut body = Vec::with_capacity(body_len);
body.extend_from_slice(&TLS_VERSION);
body.extend_from_slice(&[fill; 32]);
body.push(session_id_len as u8);
body.extend_from_slice(&[fill; 32]);
body.extend_from_slice(&(TLS_AES_128_GCM_SHA256.len() as u16).to_be_bytes());
body.extend_from_slice(&TLS_AES_128_GCM_SHA256);
body.push(1);
body.push(0);
body.extend_from_slice(&(extensions.len() as u16).to_be_bytes());
body.extend_from_slice(&extensions);
let mut handshake = Vec::with_capacity(5 + tls_len);
handshake.push(0x16);
handshake.extend_from_slice(&[0x03, 0x01]);
handshake.extend_from_slice(&(tls_len as u16).to_be_bytes());
handshake.push(0x01);
handshake.extend_from_slice(&(body_len as u32).to_be_bytes()[1..].as_ref());
handshake.extend_from_slice(&body);
handshake[tls::TLS_DIGEST_POS..tls::TLS_DIGEST_POS + tls::TLS_DIGEST_LEN].fill(0);
let digest = sha256_hmac(&SECRET, &handshake);
handshake[tls::TLS_DIGEST_POS..tls::TLS_DIGEST_POS + tls::TLS_DIGEST_LEN]
.copy_from_slice(&digest);
handshake
}
async fn wait_for_file(path: &Path) {
tokio::time::timeout(Duration::from_secs(10), async {
while !path.exists() {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("wire-test release barrier timed out");
}
async fn run_server(addr: SocketAddr, fragment_size: u16, fake_cert_len: usize) {
let socket = create_listener(
addr,
&ListenOptions {
reuse_port: false,
client_mss: Some(1400),
..Default::default()
},
)
.unwrap();
let listener = TcpListener::from_std(socket.into()).unwrap();
let (mut server, peer) = listener.accept().await.unwrap();
let mut header = [0_u8; 5];
server.read_exact(&mut header).await.unwrap();
let body_len = u16::from_be_bytes([header[3], header[4]]) as usize;
let mut client_hello = Vec::with_capacity(5 + body_len);
client_hello.extend_from_slice(&header);
client_hello.resize(5 + body_len, 0);
server.read_exact(&mut client_hello[5..]).await.unwrap();
let barrier = std::env::var("TELEMT_WIRE_BARRIER").unwrap();
let release = std::env::var("TELEMT_WIRE_RELEASE").unwrap();
std::fs::write(&barrier, b"ready").unwrap();
wait_for_file(Path::new(&release)).await;
let raw_fd = server.as_raw_fd();
let mss_before = socket2::SockRef::from(&server).tcp_mss().unwrap();
let (read_half, write_half) = server.into_split();
let mut config = ProxyConfig::default();
config.general.beobachten = false;
config.access.ignore_time_skew = true;
config.censorship.fake_cert_len = fake_cert_len;
config
.access
.users
.insert("wire".to_string(), SECRET_HEX.to_string());
let replay_checker = ReplayChecker::new(128, Duration::from_secs(60));
let rng = SecureRandom::new();
let shared = ProxySharedState::new();
let (tls_reader, mut tls_writer, user) = match
handle_tls_handshake_with_shared_and_options(
&client_hello,
read_half,
write_half,
peer,
&config,
&replay_checker,
&rng,
None,
&shared,
TlsResponseWriteOptions::tcp(raw_fd, Some(fragment_size)),
)
.await
{
HandshakeResult::Success(result) => result,
_ => panic!("wire-test FakeTLS authentication failed"),
};
assert_eq!(user, "wire");
tls_writer.write_all(&vec![0xA5; BULK_PAYLOAD_LEN]).await.unwrap();
tls_writer.shutdown().await.unwrap();
drop(tls_reader);
// SAFETY: the write half still owns the accepted socket while it is borrowed.
let borrowed_fd = unsafe { BorrowedFd::borrow_raw(raw_fd) };
let mss_after = socket2::SockRef::from(&borrowed_fd).tcp_mss().unwrap();
let metadata = std::env::var("TELEMT_WIRE_SERVER_META").unwrap();
std::fs::write(
metadata,
format!("configured_bulk_mss=1400\nmss_before={mss_before}\nmss_after={mss_after}\n"),
)
.unwrap();
}
async fn run_client(addr: SocketAddr) {
let mut client = loop {
match TcpStream::connect(addr).await {
Ok(stream) => break stream,
Err(_) => tokio::time::sleep(Duration::from_millis(20)).await,
}
};
client
.write_all(&make_valid_tls_client_hello(600))
.await
.unwrap();
let mut response = Vec::new();
client.read_to_end(&mut response).await.unwrap();
let mut offset = 0;
let mut bulk_record_start = None;
while offset < response.len() {
assert!(response.len() - offset >= 5, "truncated TLS record header");
let payload_len = u16::from_be_bytes([response[offset + 3], response[offset + 4]]) as usize;
let end = offset + 5 + payload_len;
assert!(end <= response.len(), "truncated TLS record payload");
if response[offset] == 0x17
&& payload_len == BULK_PAYLOAD_LEN
&& response[offset + 5..end].iter().all(|byte| *byte == 0xA5)
{
bulk_record_start = Some(offset);
}
offset = end;
}
let initial_response_bytes = bulk_record_start.expect("bulk TLS record missing");
let metadata = std::env::var("TELEMT_WIRE_CLIENT_META").unwrap();
std::fs::write(
metadata,
format!(
"initial_response_bytes={initial_response_bytes}\ntotal_response_bytes={}\n",
response.len()
),
)
.unwrap();
}
#[tokio::test]
#[ignore = "requires privileged netns/veth packet-capture harness"]
async fn fake_tls_fragmentation_wire_role() {
let role = std::env::var("TELEMT_WIRE_ROLE").expect("TELEMT_WIRE_ROLE is required");
let addr = std::env::var("TELEMT_WIRE_ADDR")
.unwrap_or_else(|_| "198.18.0.1:24443".to_string())
.parse()
.unwrap();
match role.as_str() {
"server" => {
let fragment_size = std::env::var("TELEMT_WIRE_FRAGMENT")
.unwrap()
.parse()
.unwrap();
let fake_cert_len = std::env::var("TELEMT_WIRE_FAKE_CERT_LEN")
.unwrap()
.parse()
.unwrap();
run_server(addr, fragment_size, fake_cert_len).await;
}
"client" => run_client(addr).await,
_ => panic!("TELEMT_WIRE_ROLE must be server or client"),
}
}
+73 -25
View File
@@ -249,7 +249,34 @@ async fn test_chunked_send_preserves_stream_and_configured_mss() {
);
assert_eq!(
mss_after, mss_before,
"chunked send must not change the configured socket MSS"
"chunked send must restore the configured socket MSS"
);
}
#[cfg(target_os = "linux")]
#[tokio::test]
async fn test_chunked_send_restores_mss_after_send_error() {
use std::os::fd::AsRawFd;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let _client = TcpStream::connect(addr).await.unwrap();
let (server, _) = listener.accept().await.unwrap();
let mss_before = socket2::SockRef::from(&server).tcp_mss().unwrap();
let shutdown_result = unsafe { libc::shutdown(server.as_raw_fd(), libc::SHUT_WR) };
assert_eq!(shutdown_result, 0);
let error = send_tcp_fragmented_fd(server.as_raw_fd(), &[0xA5; 4096], 92)
.await
.unwrap_err();
assert!(matches!(
error.kind(),
ErrorKind::BrokenPipe | ErrorKind::ConnectionReset | ErrorKind::NotConnected
));
assert_eq!(
socket2::SockRef::from(&server).tcp_mss().unwrap(),
mss_before
);
}
@@ -279,17 +306,24 @@ async fn test_chunked_send_has_no_fd_growth_after_success_and_cancellation_stres
use std::sync::Arc;
let baseline_fds = std::fs::read_dir("/proc/self/fd").unwrap().count();
let blocked_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let blocked_addr = blocked_listener.local_addr().unwrap();
let _blocked_client = TcpStream::connect(blocked_addr).await.unwrap();
let (blocked_server, _) = blocked_listener.accept().await.unwrap();
socket2::SockRef::from(&blocked_server)
.set_send_buffer_size(4 * 1024)
.unwrap();
let blocked_fd = blocked_server.as_raw_fd();
let payload = Arc::new(vec![0xA5; 1024 * 1024]);
let options = ListenOptions {
reuse_port: false,
client_mss: Some(1400),
..Default::default()
};
let blocked_socket = create_listener("127.0.0.1:0".parse().unwrap(), &options).unwrap();
let blocked_listener = TcpListener::from_std(blocked_socket.into()).unwrap();
let blocked_addr = blocked_listener.local_addr().unwrap();
for _ in 0..5_000 {
let blocked_client = TcpStream::connect(blocked_addr).await.unwrap();
let (blocked_server, _) = blocked_listener.accept().await.unwrap();
let blocked_mss_before = socket2::SockRef::from(&blocked_server).tcp_mss().unwrap();
socket2::SockRef::from(&blocked_server)
.set_send_buffer_size(4 * 1024)
.unwrap();
let blocked_fd = blocked_server.as_raw_fd();
let payload = payload.clone();
let sender = tokio::spawn(async move {
send_tcp_fragmented_fd(blocked_fd, payload.as_slice(), 92).await
@@ -297,29 +331,43 @@ async fn test_chunked_send_has_no_fd_growth_after_success_and_cancellation_stres
tokio::task::yield_now().await;
sender.abort();
let _ = sender.await;
assert_eq!(
socket2::SockRef::from(&blocked_server).tcp_mss().unwrap(),
blocked_mss_before,
"cancellation must restore the accepted socket MSS"
);
drop(blocked_server);
drop(blocked_client);
}
let success_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let success_socket = create_listener("127.0.0.1:0".parse().unwrap(), &options).unwrap();
let success_listener = TcpListener::from_std(success_socket.into()).unwrap();
let success_addr = success_listener.local_addr().unwrap();
let mut success_client = TcpStream::connect(success_addr).await.unwrap();
let (success_server, _) = success_listener.accept().await.unwrap();
let success_fd = success_server.as_raw_fd();
let reader = tokio::spawn(async move {
let mut received = vec![0_u8; 5_000];
success_client.read_exact(&mut received).await.unwrap();
received
});
for _ in 0..5_000 {
for iteration in 0..5_000 {
let mut success_client = TcpStream::connect(success_addr).await.unwrap();
let (success_server, _) = success_listener.accept().await.unwrap();
let success_mss_before = socket2::SockRef::from(&success_server).tcp_mss().unwrap();
let success_fd = success_server.as_raw_fd();
send_tcp_fragmented_fd(success_fd, &[0x5A], 92)
.await
.unwrap();
.unwrap_or_else(|error| {
panic!(
"success cycle {iteration} failed with MSS {success_mss_before}: {error}"
)
});
let mut received = [0_u8; 1];
success_client.read_exact(&mut received).await.unwrap();
assert_eq!(received, [0x5A]);
assert_eq!(
socket2::SockRef::from(&success_server).tcp_mss().unwrap(),
success_mss_before,
"successful send must restore the accepted socket MSS"
);
drop(success_server);
drop(success_client);
}
assert!(reader.await.unwrap().iter().all(|byte| *byte == 0x5A));
drop(success_server);
drop(success_listener);
drop(blocked_server);
drop(_blocked_client);
drop(blocked_listener);
let final_fds = std::fs::read_dir("/proc/self/fd").unwrap().count();