use std::io;
use std::net::{SocketAddr, TcpListener as StdListener};
use std::os::fd::{AsRawFd, FromRawFd, IntoRawFd, RawFd};
use socket2::{Domain, Protocol, Socket, Type};
pub fn tcp_listener(addr: SocketAddr, reuse_port: bool, backlog: i32) -> io::Result<StdListener> {
let domain = if addr.is_ipv4() {
Domain::IPV4
} else {
Domain::IPV6
};
let sock = Socket::new(domain, Type::STREAM, Some(Protocol::TCP))?;
sock.set_nonblocking(true)?;
sock.set_reuse_address(true)?;
if reuse_port {
set_sock_opt_bool(sock.as_raw_fd(), libc::SOL_SOCKET, libc::SO_REUSEPORT, true)?;
}
if addr.ip().is_ipv6() {
set_sock_opt_bool(
sock.as_raw_fd(),
libc::IPPROTO_IPV6,
libc::IPV6_V6ONLY,
true,
)?;
}
sock.bind(&addr.into())?;
sock.listen(backlog)?;
Ok(unsafe { StdListener::from_raw_fd(sock.into_raw_fd()) })
}
pub fn set_nodelay(fd: RawFd) -> io::Result<()> {
set_sock_opt_bool(fd, libc::IPPROTO_TCP, libc::TCP_NODELAY, true)
}
pub fn set_keepalive(fd: RawFd, idle_secs: u32) -> io::Result<()> {
set_sock_opt_bool(fd, libc::SOL_SOCKET, libc::SO_KEEPALIVE, true)?;
unsafe {
let v: libc::c_int = idle_secs as libc::c_int;
libc::setsockopt(
fd,
libc::IPPROTO_TCP,
libc::TCP_KEEPIDLE,
std::ptr::addr_of!(v).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
);
let i: libc::c_int = 3;
libc::setsockopt(
fd,
libc::IPPROTO_TCP,
libc::TCP_KEEPINTVL,
std::ptr::addr_of!(i).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
);
let c: libc::c_int = 3;
libc::setsockopt(
fd,
libc::IPPROTO_TCP,
libc::TCP_KEEPCNT,
std::ptr::addr_of!(c).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
);
}
Ok(())
}
fn set_sock_opt_bool(fd: RawFd, level: libc::c_int, name: libc::c_int, on: bool) -> io::Result<()> {
let v: libc::c_int = i32::from(on);
let rc = unsafe {
libc::setsockopt(
fd,
level,
name,
std::ptr::addr_of!(v).cast(),
std::mem::size_of::<libc::c_int>() as libc::socklen_t,
)
};
if rc == 0 {
Ok(())
} else {
Err(io::Error::last_os_error())
}
}
#[derive(Debug)]
pub struct StreamFd(pub RawFd);
impl StreamFd {
#[must_use]
pub fn fd(&self) -> RawFd {
self.0
}
}
impl Drop for StreamFd {
fn drop(&mut self) {
unsafe { libc::close(self.0) };
}
}
pub fn shutdown_write(fd: RawFd) {
unsafe {
let _ = libc::shutdown(fd, libc::SHUT_WR);
}
}
#[cfg(test)]
mod net_tests {
use super::*;
#[test]
fn tcp_listener_reuseport_binds() {
let addr: SocketAddr = "127.0.0.1:0".parse().expect("addr");
let l = tcp_listener(addr, true, 64).expect("bind");
assert!(l.local_addr().is_ok());
}
#[test]
fn set_nodelay_and_keepalive_on_socket() {
let a = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = a.local_addr().expect("addr");
let client = std::net::TcpStream::connect(addr).expect("connect");
let fd = client.as_raw_fd();
set_nodelay(fd).expect("nodelay");
set_keepalive(fd, 30).expect("keepalive");
}
#[test]
fn set_nodelay_rejects_bad_fd() {
assert!(set_nodelay(-1).is_err());
}
#[test]
fn shutdown_write_sends_fin() {
use std::io::Read as _;
let a = std::net::TcpListener::bind("127.0.0.1:0").expect("bind");
let addr = a.local_addr().expect("addr");
let client = std::net::TcpStream::connect(addr).expect("connect");
let mut server = a.incoming().next().unwrap().expect("accept");
shutdown_write(client.as_raw_fd());
let mut buf = [0u8; 1];
assert_eq!(server.read(&mut buf).expect("read"), 0);
}
}