#![allow(unsafe_code)]
use crate::error::{LinuxError, Result};
pub const MAX_BATCH_DATAGRAMS: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u32)]
pub enum SpliceFlags {
None = 0,
Move = libc::SPLICE_F_MOVE,
NonBlock = libc::SPLICE_F_NONBLOCK,
More = libc::SPLICE_F_MORE,
}
impl SpliceFlags {
#[inline]
pub fn bits(self) -> u32 {
self as u32
}
}
pub fn splice(
fd_in: i32,
off_in: Option<&mut i64>,
fd_out: i32,
off_out: Option<&mut i64>,
len: usize,
flags: u32,
) -> Result<usize> {
let safe_len = len.min(0x7FFF_FFFFusize);
let off_in_ptr = off_in.map_or(std::ptr::null_mut(), |p| p as *mut i64);
let off_out_ptr = off_out.map_or(std::ptr::null_mut(), |p| p as *mut i64);
let ret = unsafe {
libc::splice(
fd_in,
off_in_ptr,
fd_out,
off_out_ptr,
safe_len,
flags,
)
};
if ret < 0 {
Err(LinuxError::Syscall {
syscall: "splice",
errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
})
} else {
Ok(ret as usize)
}
}
pub fn sendfile(
out_fd: i32,
in_fd: i32,
offset: Option<&mut i64>,
count: usize,
) -> Result<usize> {
let off_ptr = offset.map_or(std::ptr::null_mut(), |p| p as *mut i64);
let ret = unsafe { libc::sendfile(out_fd, in_fd, off_ptr, count.min(0x7FFF_FFFFusize)) };
if ret < 0 {
Err(LinuxError::Syscall {
syscall: "sendfile",
errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
})
} else {
Ok(ret as usize)
}
}
pub fn pipe2(flags: i32) -> Result<(i32, i32)> {
let mut fds = [0i32; 2];
let ret = unsafe { libc::pipe2(fds.as_mut_ptr(), flags) };
if ret < 0 {
Err(LinuxError::Syscall {
syscall: "pipe2",
errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
})
} else {
Ok((fds[0], fds[1]))
}
}
pub fn recvmmsg(bufs: &mut [&mut [u8]], fd: i32, flags: i32) -> Result<usize> {
let count = bufs.len().min(MAX_BATCH_DATAGRAMS);
if count == 0 {
return Ok(0);
}
let mut msgs = [libc::mmsghdr {
msg_hdr: libc::msghdr {
msg_name: std::ptr::null_mut(),
msg_namelen: 0,
msg_iov: std::ptr::null_mut(),
msg_iovlen: 0,
msg_control: std::ptr::null_mut(),
msg_controllen: 0,
msg_flags: 0,
},
msg_len: 0,
}; MAX_BATCH_DATAGRAMS];
let mut iovs = [libc::iovec {
iov_base: std::ptr::null_mut(),
iov_len: 0,
}; MAX_BATCH_DATAGRAMS];
for i in 0..count {
iovs[i] = libc::iovec {
iov_base: bufs[i].as_mut_ptr() as *mut std::ffi::c_void,
iov_len: bufs[i].len(),
};
msgs[i].msg_hdr.msg_iov = &mut iovs[i];
msgs[i].msg_hdr.msg_iovlen = 1;
}
let ret = unsafe {
libc::recvmmsg(
fd,
msgs.as_mut_ptr(),
count as u32,
flags,
std::ptr::null_mut(),
)
};
if ret < 0 {
Err(LinuxError::Syscall {
syscall: "recvmmsg",
errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
})
} else {
Ok(ret as usize)
}
}
pub fn sendmmsg(bufs: &[&[u8]], fd: i32, dest: &libc::sockaddr_storage, flags: i32) -> Result<usize> {
let count = bufs.len().min(MAX_BATCH_DATAGRAMS);
if count == 0 {
return Ok(0);
}
let mut msgs = [libc::mmsghdr {
msg_hdr: libc::msghdr {
msg_name: std::ptr::null_mut(),
msg_namelen: 0,
msg_iov: std::ptr::null_mut(),
msg_iovlen: 0,
msg_control: std::ptr::null_mut(),
msg_controllen: 0,
msg_flags: 0,
},
msg_len: 0,
}; MAX_BATCH_DATAGRAMS];
let mut iovs = [libc::iovec {
iov_base: std::ptr::null_mut(),
iov_len: 0,
}; MAX_BATCH_DATAGRAMS];
let dest_ptr = dest as *const libc::sockaddr_storage as *mut std::ffi::c_void;
let dest_len = std::mem::size_of::<libc::sockaddr_storage>() as u32;
for i in 0..count {
iovs[i] = libc::iovec {
iov_base: bufs[i].as_ptr() as *mut std::ffi::c_void,
iov_len: bufs[i].len(),
};
msgs[i].msg_hdr.msg_name = dest_ptr;
msgs[i].msg_hdr.msg_namelen = dest_len;
msgs[i].msg_hdr.msg_iov = &mut iovs[i];
msgs[i].msg_hdr.msg_iovlen = 1;
}
let ret = unsafe {
libc::sendmmsg(
fd,
msgs.as_mut_ptr(),
count as u32,
flags,
)
};
if ret < 0 {
Err(LinuxError::Syscall {
syscall: "sendmmsg",
errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
})
} else {
Ok(ret as usize)
}
}
pub fn splice_bidirectional(
client_fd: i32,
upstream_fd: i32,
pipe_buf_size: usize,
) -> Result<(usize, usize)> {
if pipe_buf_size == 0 {
return Err(LinuxError::InsufficientResources(
"pipe_buf_size 不能为 0(会导致无限 0 字节 splice 调用)".to_string(),
));
}
let (c2u_read, c2u_write) = pipe2(0)?;
let (u2c_read, u2c_write) = pipe2(0)?;
let c2u_read_guard = FdGuard::new(c2u_read);
let mut c2u_write_guard = Some(FdGuard::new(c2u_write));
let u2c_read_guard = FdGuard::new(u2c_read);
let mut u2c_write_guard = Some(FdGuard::new(u2c_write));
let relay = |src_fd: i32, pipe_r: i32, pipe_w: i32, dst_fd: i32, p_size: usize| -> std::result::Result<usize, LinuxError> {
let mut total = 0usize;
loop {
let n = match splice(src_fd, None, pipe_w, None, p_size, libc::SPLICE_F_MOVE) {
Ok(n) => n,
Err(LinuxError::Syscall { syscall: _, errno }) if errno == libc::EPIPE => {
unsafe { libc::close(pipe_w); }
return Ok(total);
}
Err(e) => {
unsafe { libc::close(pipe_w); }
return Err(e);
}
};
if n == 0 {
unsafe { libc::close(pipe_w); }
loop {
let m = splice(pipe_r, None, dst_fd, None, p_size, libc::SPLICE_F_MOVE)?;
if m == 0 {
break;
}
total = total.checked_add(m).ok_or_else(|| LinuxError::InsufficientResources(
"splice byte count overflow".to_string()
))?;
}
unsafe { let _ = libc::shutdown(dst_fd, libc::SHUT_WR); }
return Ok(total);
}
total = total.checked_add(n).ok_or_else(|| LinuxError::InsufficientResources(
"splice byte count overflow".to_string()
))?;
let m = splice(pipe_r, None, dst_fd, None, n, libc::SPLICE_F_MOVE)?;
debug_assert_eq!(m, n, "splice through pipe must preserve byte count");
}
};
let t1 = match std::thread::Builder::new()
.name("zenith-splice-c2u".to_string())
.spawn(move || relay(client_fd, c2u_read, c2u_write, upstream_fd, pipe_buf_size))
{
Ok(t) => {
if let Some(g) = c2u_write_guard.take() {
let _ = g.into_raw();
}
t
}
Err(e) => {
return Err(LinuxError::InsufficientResources(format!(
"spawn c2u thread failed: {e}"
)));
}
};
let t2 = match std::thread::Builder::new()
.name("zenith-splice-u2c".to_string())
.spawn(move || relay(upstream_fd, u2c_read, u2c_write, client_fd, pipe_buf_size))
{
Ok(t) => {
if let Some(g) = u2c_write_guard.take() {
let _ = g.into_raw();
}
t
}
Err(e) => {
unsafe {
let _ = libc::shutdown(client_fd, libc::SHUT_RDWR);
let _ = libc::shutdown(upstream_fd, libc::SHUT_RDWR);
}
let _ = t1.join();
return Err(LinuxError::InsufficientResources(format!(
"spawn u2c thread failed: {e}"
)));
}
};
let r1 = t1.join().map_err(|_| LinuxError::InsufficientResources(
"c2u relay thread panicked".to_string()
))??;
let r2 = t2.join().map_err(|_| LinuxError::InsufficientResources(
"u2c relay thread panicked".to_string()
))??;
let _ = (c2u_read_guard, u2c_read_guard);
Ok((r1, r2))
}
#[derive(Debug)]
pub struct FdGuard {
fd: i32,
closed: bool,
}
impl FdGuard {
#[inline]
pub fn new(fd: i32) -> Self {
Self { fd, closed: false }
}
#[inline]
pub fn as_raw(&self) -> i32 {
self.fd
}
#[inline]
pub fn into_raw(mut self) -> i32 {
self.closed = true;
self.fd
}
}
impl Drop for FdGuard {
fn drop(&mut self) {
if !self.closed && self.fd >= 0 {
unsafe { let _ = libc::close(self.fd); }
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_splice_flags_bits() {
assert_eq!(SpliceFlags::None.bits(), 0);
assert_eq!(SpliceFlags::Move.bits(), libc::SPLICE_F_MOVE);
assert_eq!(SpliceFlags::NonBlock.bits(), libc::SPLICE_F_NONBLOCK);
assert_eq!(SpliceFlags::More.bits(), libc::SPLICE_F_MORE);
}
#[test]
fn test_pipe2_create() {
let (r, w) = pipe2(libc::O_NONBLOCK).expect("pipe2");
assert!(r >= 0);
assert!(w >= 0);
assert_ne!(r, w);
unsafe {
libc::close(r);
libc::close(w);
}
}
#[test]
fn test_pipe2_invalid_flags() {
let result = pipe2(0x7FFF_FFFF);
let _ = result;
}
#[test]
fn test_splice_eof_pipe() {
let (r, w) = pipe2(0).expect("pipe2");
unsafe { libc::close(w); }
let n = splice(r, None, -1, None, 1024, 0);
assert!(n.is_err());
unsafe { libc::close(r); }
}
#[test]
fn test_fd_guard_closes() {
let (r, w) = pipe2(0).expect("pipe2");
unsafe { libc::close(w); }
{
let _guard = FdGuard::new(r);
}
}
#[test]
fn test_fd_guard_into_raw() {
let (r, w) = pipe2(0).expect("pipe2");
unsafe { libc::close(w); }
let guard = FdGuard::new(r);
let raw = guard.into_raw();
assert_eq!(raw, r);
unsafe { libc::close(r); }
}
#[test]
fn test_recvmmsg_empty_bufs() {
let mut bufs: [&mut [u8]; 0] = [];
let result = recvmmsg(&mut bufs, -1, 0);
assert_eq!(result.unwrap(), 0);
}
#[test]
fn test_sendmmsg_empty_bufs() {
let bufs: [&[u8]; 0] = [];
let dest: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
let result = sendmmsg(&bufs, -1, &dest, 0);
assert_eq!(result.unwrap(), 0);
}
#[test]
fn test_sendfile_invalid_fd() {
let result = sendfile(-1, -1, None, 1024);
assert!(result.is_err());
let err = result.unwrap_err();
match err {
LinuxError::Syscall { syscall, .. } => assert_eq!(syscall, "sendfile"),
_ => panic!("expected Syscall error"),
}
}
#[test]
fn test_splice_invalid_fd() {
let result = splice(-1, None, -1, None, 1024, 0);
assert!(result.is_err());
let err = result.unwrap_err();
match err {
LinuxError::Syscall { syscall, .. } => assert_eq!(syscall, "splice"),
_ => panic!("expected Syscall error"),
}
}
#[test]
fn test_max_batch_datagrams_constant() {
assert_eq!(MAX_BATCH_DATAGRAMS, 64);
const { assert!(MAX_BATCH_DATAGRAMS > 0) };
}
fn open_pipe_inodes() -> std::collections::HashSet<String> {
let mut set = std::collections::HashSet::new();
if let Ok(dir) = std::fs::read_dir("/proc/self/fd") {
for entry in dir.flatten() {
if let Ok(target) = std::fs::read_link(entry.path()) {
let s = target.to_string_lossy().into_owned();
if s.starts_with("pipe:[") {
set.insert(s);
}
}
}
}
set
}
#[test]
fn test_splice_bidirectional_tcp_loopback() {
use std::io::{Read, Write};
use std::net::{Shutdown, TcpListener, TcpStream};
use std::os::unix::io::AsRawFd;
let pipe_baseline = open_pipe_inodes();
let echo_listener = match TcpListener::bind("127.0.0.1:0") {
Ok(l) => l,
Err(_) => return, };
let echo_addr = echo_listener.local_addr().unwrap();
let echo_thread = std::thread::spawn(move || -> usize {
let (mut s, _) = echo_listener.accept().unwrap();
let mut buf = [0u8; 4096];
let mut echoed = 0usize;
loop {
match s.read(&mut buf) {
Ok(0) | Err(_) => break, Ok(n) => {
s.write_all(&buf[..n]).unwrap();
echoed += n;
}
}
}
let _ = s.shutdown(Shutdown::Write); echoed
});
let proxy_listener = TcpListener::bind("127.0.0.1:0").unwrap();
let proxy_addr = proxy_listener.local_addr().unwrap();
let upstream = TcpStream::connect(echo_addr).unwrap();
let mut client = TcpStream::connect(proxy_addr).unwrap();
let (proxy_side, _) = proxy_listener.accept().unwrap();
let proxy_fd = proxy_side.as_raw_fd();
let upstream_fd = upstream.as_raw_fd();
let relay_thread = std::thread::spawn(move || {
splice_bidirectional(proxy_fd, upstream_fd, 16384)
});
let payload: Vec<u8> = (0..65536u32).map(|i| (i % 251) as u8).collect();
client.write_all(&payload).unwrap();
client.shutdown(Shutdown::Write).unwrap();
let mut got = Vec::with_capacity(payload.len());
let mut tmp = [0u8; 8192];
loop {
match client.read(&mut tmp) {
Ok(0) => break, Ok(n) => got.extend_from_slice(&tmp[..n]),
Err(e) => panic!("client read failed: {e}"),
}
}
assert_eq!(got.len(), payload.len(), "echo 字节数必须守恒");
assert!(got == payload, "echo 内容必须逐字节一致");
let echoed = echo_thread.join().unwrap();
assert_eq!(echoed, payload.len(), "echo 服务端必须读满全部负载");
let (c2u, u2c) = relay_thread.join().unwrap().unwrap();
assert_eq!(c2u, payload.len(), "c2u 方向字节数必须守恒");
assert_eq!(u2c, payload.len(), "u2c 方向字节数必须守恒");
drop(client);
drop(proxy_side);
drop(upstream);
drop(proxy_listener);
let mut extra = Vec::new();
for _ in 0..100 {
extra = open_pipe_inodes()
.difference(&pipe_baseline)
.cloned()
.collect::<Vec<_>>();
if extra.is_empty() {
break;
}
std::thread::sleep(std::time::Duration::from_millis(10));
}
assert!(extra.is_empty(), "splice_bidirectional 泄漏 pipe fd: {extra:?}");
}
}