1#![allow(unsafe_code)]
24
25use crate::error::{LinuxError, Result};
26
27pub const MAX_BATCH_DATAGRAMS: usize = 64;
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32#[repr(u32)]
33pub enum SpliceFlags {
34 None = 0,
36 Move = libc::SPLICE_F_MOVE,
38 NonBlock = libc::SPLICE_F_NONBLOCK,
40 More = libc::SPLICE_F_MORE,
42}
43
44impl SpliceFlags {
45 #[inline]
47 pub fn bits(self) -> u32 {
48 self as u32
49 }
50}
51
52pub fn splice(
70 fd_in: i32,
71 off_in: Option<&mut i64>,
72 fd_out: i32,
73 off_out: Option<&mut i64>,
74 len: usize,
75 flags: u32,
76) -> Result<usize> {
77 let safe_len = len.min(0x7FFF_FFFFusize);
79
80 let off_in_ptr = off_in.map_or(std::ptr::null_mut(), |p| p as *mut i64);
81 let off_out_ptr = off_out.map_or(std::ptr::null_mut(), |p| p as *mut i64);
82
83 let ret = unsafe {
87 libc::splice(
88 fd_in,
89 off_in_ptr,
90 fd_out,
91 off_out_ptr,
92 safe_len,
93 flags,
94 )
95 };
96
97 if ret < 0 {
98 Err(LinuxError::Syscall {
99 syscall: "splice",
100 errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
101 })
102 } else {
103 Ok(ret as usize)
104 }
105}
106
107pub fn sendfile(
120 out_fd: i32,
121 in_fd: i32,
122 offset: Option<&mut i64>,
123 count: usize,
124) -> Result<usize> {
125 let off_ptr = offset.map_or(std::ptr::null_mut(), |p| p as *mut i64);
126
127 let ret = unsafe { libc::sendfile(out_fd, in_fd, off_ptr, count.min(0x7FFF_FFFFusize)) };
130
131 if ret < 0 {
132 Err(LinuxError::Syscall {
133 syscall: "sendfile",
134 errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
135 })
136 } else {
137 Ok(ret as usize)
138 }
139}
140
141pub fn pipe2(flags: i32) -> Result<(i32, i32)> {
149 let mut fds = [0i32; 2];
150
151 let ret = unsafe { libc::pipe2(fds.as_mut_ptr(), flags) };
153
154 if ret < 0 {
155 Err(LinuxError::Syscall {
156 syscall: "pipe2",
157 errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
158 })
159 } else {
160 Ok((fds[0], fds[1]))
161 }
162}
163
164pub fn recvmmsg(bufs: &mut [&mut [u8]], fd: i32, flags: i32) -> Result<usize> {
176 let count = bufs.len().min(MAX_BATCH_DATAGRAMS);
177 if count == 0 {
178 return Ok(0);
179 }
180
181 let mut msgs = [libc::mmsghdr {
183 msg_hdr: libc::msghdr {
184 msg_name: std::ptr::null_mut(),
185 msg_namelen: 0,
186 msg_iov: std::ptr::null_mut(),
187 msg_iovlen: 0,
188 msg_control: std::ptr::null_mut(),
189 msg_controllen: 0,
190 msg_flags: 0,
191 },
192 msg_len: 0,
193 }; MAX_BATCH_DATAGRAMS];
194
195 let mut iovs = [libc::iovec {
197 iov_base: std::ptr::null_mut(),
198 iov_len: 0,
199 }; MAX_BATCH_DATAGRAMS];
200
201 for i in 0..count {
203 iovs[i] = libc::iovec {
204 iov_base: bufs[i].as_mut_ptr() as *mut std::ffi::c_void,
205 iov_len: bufs[i].len(),
206 };
207 msgs[i].msg_hdr.msg_iov = &mut iovs[i];
208 msgs[i].msg_hdr.msg_iovlen = 1;
209 }
210
211 let ret = unsafe {
215 libc::recvmmsg(
216 fd,
217 msgs.as_mut_ptr(),
218 count as u32,
219 flags,
220 std::ptr::null_mut(),
221 )
222 };
223
224 if ret < 0 {
225 Err(LinuxError::Syscall {
226 syscall: "recvmmsg",
227 errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
228 })
229 } else {
230 Ok(ret as usize)
231 }
232}
233
234pub fn sendmmsg(bufs: &[&[u8]], fd: i32, dest: &libc::sockaddr_storage, flags: i32) -> Result<usize> {
247 let count = bufs.len().min(MAX_BATCH_DATAGRAMS);
248 if count == 0 {
249 return Ok(0);
250 }
251
252 let mut msgs = [libc::mmsghdr {
253 msg_hdr: libc::msghdr {
254 msg_name: std::ptr::null_mut(),
255 msg_namelen: 0,
256 msg_iov: std::ptr::null_mut(),
257 msg_iovlen: 0,
258 msg_control: std::ptr::null_mut(),
259 msg_controllen: 0,
260 msg_flags: 0,
261 },
262 msg_len: 0,
263 }; MAX_BATCH_DATAGRAMS];
264
265 let mut iovs = [libc::iovec {
266 iov_base: std::ptr::null_mut(),
267 iov_len: 0,
268 }; MAX_BATCH_DATAGRAMS];
269
270 let dest_ptr = dest as *const libc::sockaddr_storage as *mut std::ffi::c_void;
272 let dest_len = std::mem::size_of::<libc::sockaddr_storage>() as u32;
273
274 for i in 0..count {
275 iovs[i] = libc::iovec {
276 iov_base: bufs[i].as_ptr() as *mut std::ffi::c_void,
277 iov_len: bufs[i].len(),
278 };
279 msgs[i].msg_hdr.msg_name = dest_ptr;
280 msgs[i].msg_hdr.msg_namelen = dest_len;
281 msgs[i].msg_hdr.msg_iov = &mut iovs[i];
282 msgs[i].msg_hdr.msg_iovlen = 1;
283 }
284
285 let ret = unsafe {
289 libc::sendmmsg(
290 fd,
291 msgs.as_mut_ptr(),
292 count as u32,
293 flags,
294 )
295 };
296
297 if ret < 0 {
298 Err(LinuxError::Syscall {
299 syscall: "sendmmsg",
300 errno: std::io::Error::last_os_error().raw_os_error().unwrap_or(0),
301 })
302 } else {
303 Ok(ret as usize)
304 }
305}
306
307pub fn splice_bidirectional(
328 client_fd: i32,
329 upstream_fd: i32,
330 pipe_buf_size: usize,
331) -> Result<(usize, usize)> {
332 if pipe_buf_size == 0 {
334 return Err(LinuxError::InsufficientResources(
335 "pipe_buf_size 不能为 0(会导致无限 0 字节 splice 调用)".to_string(),
336 ));
337 }
338
339 let (c2u_read, c2u_write) = pipe2(0)?;
341 let (u2c_read, u2c_write) = pipe2(0)?;
342 let c2u_read_guard = FdGuard::new(c2u_read);
346 let mut c2u_write_guard = Some(FdGuard::new(c2u_write));
347 let u2c_read_guard = FdGuard::new(u2c_read);
348 let mut u2c_write_guard = Some(FdGuard::new(u2c_write));
349
350 let relay = |src_fd: i32, pipe_r: i32, pipe_w: i32, dst_fd: i32, p_size: usize| -> std::result::Result<usize, LinuxError> {
353 let mut total = 0usize;
354 loop {
355 let n = match splice(src_fd, None, pipe_w, None, p_size, libc::SPLICE_F_MOVE) {
357 Ok(n) => n,
358 Err(LinuxError::Syscall { syscall: _, errno }) if errno == libc::EPIPE => {
359 unsafe { libc::close(pipe_w); }
363 return Ok(total);
364 }
365 Err(e) => {
366 unsafe { libc::close(pipe_w); }
369 return Err(e);
370 }
371 };
372 if n == 0 {
373 unsafe { libc::close(pipe_w); }
376 loop {
378 let m = splice(pipe_r, None, dst_fd, None, p_size, libc::SPLICE_F_MOVE)?;
379 if m == 0 {
380 break;
381 }
382 total = total.checked_add(m).ok_or_else(|| LinuxError::InsufficientResources(
383 "splice byte count overflow".to_string()
384 ))?;
385 }
386 unsafe { let _ = libc::shutdown(dst_fd, libc::SHUT_WR); }
389 return Ok(total);
390 }
391 total = total.checked_add(n).ok_or_else(|| LinuxError::InsufficientResources(
392 "splice byte count overflow".to_string()
393 ))?;
394 let m = splice(pipe_r, None, dst_fd, None, n, libc::SPLICE_F_MOVE)?;
396 debug_assert_eq!(m, n, "splice through pipe must preserve byte count");
398 }
399 };
400
401 let t1 = match std::thread::Builder::new()
403 .name("zenith-splice-c2u".to_string())
404 .spawn(move || relay(client_fd, c2u_read, c2u_write, upstream_fd, pipe_buf_size))
405 {
406 Ok(t) => {
407 if let Some(g) = c2u_write_guard.take() {
408 let _ = g.into_raw();
409 }
410 t
411 }
412 Err(e) => {
413 return Err(LinuxError::InsufficientResources(format!(
415 "spawn c2u thread failed: {e}"
416 )));
417 }
418 };
419
420 let t2 = match std::thread::Builder::new()
421 .name("zenith-splice-u2c".to_string())
422 .spawn(move || relay(upstream_fd, u2c_read, u2c_write, client_fd, pipe_buf_size))
423 {
424 Ok(t) => {
425 if let Some(g) = u2c_write_guard.take() {
426 let _ = g.into_raw();
427 }
428 t
429 }
430 Err(e) => {
431 unsafe {
436 let _ = libc::shutdown(client_fd, libc::SHUT_RDWR);
437 let _ = libc::shutdown(upstream_fd, libc::SHUT_RDWR);
438 }
439 let _ = t1.join();
440 return Err(LinuxError::InsufficientResources(format!(
441 "spawn u2c thread failed: {e}"
442 )));
443 }
444 };
445
446 let r1 = t1.join().map_err(|_| LinuxError::InsufficientResources(
447 "c2u relay thread panicked".to_string()
448 ))??;
449 let r2 = t2.join().map_err(|_| LinuxError::InsufficientResources(
450 "u2c relay thread panicked".to_string()
451 ))??;
452
453 let _ = (c2u_read_guard, u2c_read_guard);
456
457 Ok((r1, r2))
459}
460
461#[derive(Debug)]
463pub struct FdGuard {
464 fd: i32,
465 closed: bool,
466}
467
468impl FdGuard {
469 #[inline]
471 pub fn new(fd: i32) -> Self {
472 Self { fd, closed: false }
473 }
474
475 #[inline]
477 pub fn as_raw(&self) -> i32 {
478 self.fd
479 }
480
481 #[inline]
483 pub fn into_raw(mut self) -> i32 {
484 self.closed = true;
485 self.fd
486 }
487}
488
489impl Drop for FdGuard {
490 fn drop(&mut self) {
491 if !self.closed && self.fd >= 0 {
492 unsafe { let _ = libc::close(self.fd); }
494 }
495 }
496}
497
498#[cfg(test)]
499mod tests {
500 use super::*;
501
502 #[test]
503 fn test_splice_flags_bits() {
504 assert_eq!(SpliceFlags::None.bits(), 0);
505 assert_eq!(SpliceFlags::Move.bits(), libc::SPLICE_F_MOVE);
506 assert_eq!(SpliceFlags::NonBlock.bits(), libc::SPLICE_F_NONBLOCK);
507 assert_eq!(SpliceFlags::More.bits(), libc::SPLICE_F_MORE);
508 }
509
510 #[test]
511 fn test_pipe2_create() {
512 let (r, w) = pipe2(libc::O_NONBLOCK).expect("pipe2");
514 assert!(r >= 0);
515 assert!(w >= 0);
516 assert_ne!(r, w);
517 unsafe {
519 libc::close(r);
520 libc::close(w);
521 }
522 }
523
524 #[test]
525 fn test_pipe2_invalid_flags() {
526 let result = pipe2(0x7FFF_FFFF);
528 let _ = result;
530 }
531
532 #[test]
533 fn test_splice_eof_pipe() {
534 let (r, w) = pipe2(0).expect("pipe2");
536 unsafe { libc::close(w); }
538 let n = splice(r, None, -1, None, 1024, 0);
539 assert!(n.is_err());
541 unsafe { libc::close(r); }
542 }
543
544 #[test]
545 fn test_fd_guard_closes() {
546 let (r, w) = pipe2(0).expect("pipe2");
547 unsafe { libc::close(w); }
548 {
549 let _guard = FdGuard::new(r);
550 }
552 }
554
555 #[test]
556 fn test_fd_guard_into_raw() {
557 let (r, w) = pipe2(0).expect("pipe2");
558 unsafe { libc::close(w); }
559 let guard = FdGuard::new(r);
560 let raw = guard.into_raw();
561 assert_eq!(raw, r);
562 unsafe { libc::close(r); }
564 }
565
566 #[test]
567 fn test_recvmmsg_empty_bufs() {
568 let mut bufs: [&mut [u8]; 0] = [];
570 let result = recvmmsg(&mut bufs, -1, 0);
571 assert_eq!(result.unwrap(), 0);
572 }
573
574 #[test]
575 fn test_sendmmsg_empty_bufs() {
576 let bufs: [&[u8]; 0] = [];
577 let dest: libc::sockaddr_storage = unsafe { std::mem::zeroed() };
579 let result = sendmmsg(&bufs, -1, &dest, 0);
580 assert_eq!(result.unwrap(), 0);
581 }
582
583 #[test]
584 fn test_sendfile_invalid_fd() {
585 let result = sendfile(-1, -1, None, 1024);
586 assert!(result.is_err());
587 let err = result.unwrap_err();
588 match err {
589 LinuxError::Syscall { syscall, .. } => assert_eq!(syscall, "sendfile"),
590 _ => panic!("expected Syscall error"),
591 }
592 }
593
594 #[test]
595 fn test_splice_invalid_fd() {
596 let result = splice(-1, None, -1, None, 1024, 0);
597 assert!(result.is_err());
598 let err = result.unwrap_err();
599 match err {
600 LinuxError::Syscall { syscall, .. } => assert_eq!(syscall, "splice"),
601 _ => panic!("expected Syscall error"),
602 }
603 }
604
605 #[test]
606 fn test_max_batch_datagrams_constant() {
607 assert_eq!(MAX_BATCH_DATAGRAMS, 64);
609 const { assert!(MAX_BATCH_DATAGRAMS > 0) };
610 }
611
612 fn open_pipe_inodes() -> std::collections::HashSet<String> {
614 let mut set = std::collections::HashSet::new();
615 if let Ok(dir) = std::fs::read_dir("/proc/self/fd") {
616 for entry in dir.flatten() {
617 if let Ok(target) = std::fs::read_link(entry.path()) {
618 let s = target.to_string_lossy().into_owned();
619 if s.starts_with("pipe:[") {
620 set.insert(s);
621 }
622 }
623 }
624 }
625 set
626 }
627
628 #[test]
637 fn test_splice_bidirectional_tcp_loopback() {
638 use std::io::{Read, Write};
639 use std::net::{Shutdown, TcpListener, TcpStream};
640 use std::os::unix::io::AsRawFd;
641
642 let pipe_baseline = open_pipe_inodes();
643
644 let echo_listener = match TcpListener::bind("127.0.0.1:0") {
646 Ok(l) => l,
647 Err(_) => return, };
649 let echo_addr = echo_listener.local_addr().unwrap();
650 let echo_thread = std::thread::spawn(move || -> usize {
651 let (mut s, _) = echo_listener.accept().unwrap();
652 let mut buf = [0u8; 4096];
653 let mut echoed = 0usize;
654 loop {
655 match s.read(&mut buf) {
656 Ok(0) | Err(_) => break, Ok(n) => {
658 s.write_all(&buf[..n]).unwrap();
659 echoed += n;
660 }
661 }
662 }
663 let _ = s.shutdown(Shutdown::Write); echoed
665 });
666
667 let proxy_listener = TcpListener::bind("127.0.0.1:0").unwrap();
669 let proxy_addr = proxy_listener.local_addr().unwrap();
670 let upstream = TcpStream::connect(echo_addr).unwrap();
671 let mut client = TcpStream::connect(proxy_addr).unwrap();
672 let (proxy_side, _) = proxy_listener.accept().unwrap();
673
674 let proxy_fd = proxy_side.as_raw_fd();
675 let upstream_fd = upstream.as_raw_fd();
676 let relay_thread = std::thread::spawn(move || {
677 splice_bidirectional(proxy_fd, upstream_fd, 16384)
678 });
679
680 let payload: Vec<u8> = (0..65536u32).map(|i| (i % 251) as u8).collect();
682 client.write_all(&payload).unwrap();
683 client.shutdown(Shutdown::Write).unwrap();
684
685 let mut got = Vec::with_capacity(payload.len());
687 let mut tmp = [0u8; 8192];
688 loop {
689 match client.read(&mut tmp) {
690 Ok(0) => break, Ok(n) => got.extend_from_slice(&tmp[..n]),
692 Err(e) => panic!("client read failed: {e}"),
693 }
694 }
695 assert_eq!(got.len(), payload.len(), "echo 字节数必须守恒");
696 assert!(got == payload, "echo 内容必须逐字节一致");
697
698 let echoed = echo_thread.join().unwrap();
700 assert_eq!(echoed, payload.len(), "echo 服务端必须读满全部负载");
701
702 let (c2u, u2c) = relay_thread.join().unwrap().unwrap();
704 assert_eq!(c2u, payload.len(), "c2u 方向字节数必须守恒");
705 assert_eq!(u2c, payload.len(), "u2c 方向字节数必须守恒");
706
707 drop(client);
709 drop(proxy_side);
710 drop(upstream);
711 drop(proxy_listener);
712
713 let mut extra = Vec::new();
717 for _ in 0..100 {
718 extra = open_pipe_inodes()
719 .difference(&pipe_baseline)
720 .cloned()
721 .collect::<Vec<_>>();
722 if extra.is_empty() {
723 break;
724 }
725 std::thread::sleep(std::time::Duration::from_millis(10));
726 }
727 assert!(extra.is_empty(), "splice_bidirectional 泄漏 pipe fd: {extra:?}");
728 }
729}