use core::fmt;
use crate::{AddressFamily, NetAddr};
pub const BIND_TABLE_CAPACITY: usize = 64;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AcceptError {
BindFailed(String),
AcceptFailed(String),
ListenTableFull,
EndpointNotFound,
AddrInUse,
AddrNotAvailable,
Io(String),
}
impl fmt::Display for AcceptError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
AcceptError::BindFailed(msg) => write!(f, "bind failed: {}", msg),
AcceptError::AcceptFailed(msg) => write!(f, "accept failed: {}", msg),
AcceptError::ListenTableFull => {
write!(f, "listen table full: capacity={}", BIND_TABLE_CAPACITY)
}
AcceptError::EndpointNotFound => write!(f, "endpoint not found in bind table"),
AcceptError::AddrInUse => write!(f, "address already in use"),
AcceptError::AddrNotAvailable => write!(f, "address not available"),
AcceptError::Io(msg) => write!(f, "io error: {}", msg),
}
}
}
impl std::error::Error for AcceptError {}
impl From<std::io::Error> for AcceptError {
fn from(err: std::io::Error) -> Self {
match err.kind() {
std::io::ErrorKind::AddrInUse => AcceptError::AddrInUse,
std::io::ErrorKind::AddrNotAvailable => AcceptError::AddrNotAvailable,
_ => AcceptError::Io(err.to_string()),
}
}
}
pub trait Accept {
type Conn;
type Error;
fn try_accept(&mut self) -> Result<Option<AcceptedConn<Self::Conn>>, Self::Error>;
}
#[derive(Debug, Clone)]
pub struct AcceptedConn<C> {
pub conn: C,
pub remote: NetAddr,
pub local: NetAddr,
pub conn_id: u64,
}
impl<C> AcceptedConn<C> {
#[inline]
pub const fn new(conn: C, remote: NetAddr, local: NetAddr, conn_id: u64) -> Self {
Self {
conn,
remote,
local,
conn_id,
}
}
#[inline]
pub fn into_parts(self) -> (C, NetAddr, NetAddr, u64) {
(self.conn, self.remote, self.local, self.conn_id)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ListenEndpoint {
pub listen_id: u64,
pub addr: NetAddr,
pub backlog: u32,
pub active_syn_count: u32,
}
#[derive(Debug, Clone)]
pub struct BindTable {
endpoints: Vec<Option<ListenEndpoint>>,
free_list: Vec<usize>,
count: u8,
capacity_: usize,
next_id: u64,
}
impl Default for BindTable {
#[inline]
fn default() -> Self {
Self::new()
}
}
impl BindTable {
#[inline]
pub fn new() -> Self {
Self::with_capacity(BIND_TABLE_CAPACITY)
}
#[inline]
pub fn with_capacity(capacity: usize) -> Self {
Self {
endpoints: Vec::with_capacity(capacity),
free_list: Vec::new(),
count: 0,
capacity_: capacity,
next_id: 1,
}
}
#[inline]
pub fn capacity(&self) -> usize {
self.capacity_
}
#[inline]
pub fn len(&self) -> usize {
self.count as usize
}
#[inline]
pub fn is_empty(&self) -> bool {
self.count == 0
}
pub fn register(&mut self, addr: NetAddr, backlog: u32) -> Result<u64, AcceptError> {
if self.is_listening(addr) {
return Err(AcceptError::AddrInUse);
}
if self.count as usize >= self.capacity_ {
return Err(AcceptError::ListenTableFull);
}
let listen_id = self.next_id;
self.next_id = self
.next_id
.checked_add(1)
.ok_or(AcceptError::ListenTableFull)?;
let endpoint = ListenEndpoint {
listen_id,
addr,
backlog,
active_syn_count: 0,
};
if let Some(idx) = self.free_list.pop() {
self.endpoints[idx] = Some(endpoint);
} else {
self.endpoints.push(Some(endpoint));
}
self.count = self
.count
.checked_add(1)
.ok_or(AcceptError::ListenTableFull)?;
Ok(listen_id)
}
#[inline]
pub fn is_listening(&self, addr: NetAddr) -> bool {
self.endpoints.iter().any(|e| match e {
Some(ep) => ep.addr == addr,
None => false,
})
}
#[inline]
pub fn get_by_id(&self, listen_id: u64) -> Option<&ListenEndpoint> {
self.endpoints.iter().find_map(|e| match e {
Some(ep) if ep.listen_id == listen_id => Some(ep),
_ => None,
})
}
#[inline]
pub fn get_by_id_mut(&mut self, listen_id: u64) -> Option<&mut ListenEndpoint> {
self.endpoints.iter_mut().find_map(|e| match e {
Some(ep) if ep.listen_id == listen_id => Some(ep),
_ => None,
})
}
pub fn deregister(&mut self, listen_id: u64) {
let found = self
.endpoints
.iter_mut()
.enumerate()
.find(|(_, e)| matches!(e, Some(ep) if ep.listen_id == listen_id));
if let Some((idx, slot)) = found {
*slot = None;
self.free_list.push(idx);
self.count = self.count.saturating_sub(1);
}
}
#[inline]
pub fn iter(&self) -> impl Iterator<Item = &ListenEndpoint> {
self.endpoints.iter().filter_map(|e| e.as_ref())
}
}
#[inline]
pub fn netaddr_to_socketaddr(addr: &NetAddr) -> Option<std::net::SocketAddr> {
match addr.family() {
AddressFamily::Ipv4 => {
let bytes = addr.ipv4_bytes();
let ip = std::net::Ipv4Addr::from(bytes);
Some(std::net::SocketAddr::V4(std::net::SocketAddrV4::new(
ip,
addr.port(),
)))
}
AddressFamily::Ipv6 => {
let bytes = addr.ipv6_bytes();
let ip = std::net::Ipv6Addr::from(bytes);
Some(std::net::SocketAddr::V6(std::net::SocketAddrV6::new(
ip,
addr.port(),
0,
0,
)))
}
}
}
#[inline]
pub fn socketaddr_to_netaddr(addr: &std::net::SocketAddr) -> NetAddr {
match addr {
std::net::SocketAddr::V4(v4) => NetAddr::new_ipv4(v4.ip().octets(), v4.port()),
std::net::SocketAddr::V6(v6) => {
let octets = v6.ip().octets();
NetAddr::new_ipv6(octets, v6.port())
}
}
}
#[inline]
pub fn register_tcp_listener(
table: &mut BindTable,
addr: NetAddr,
backlog: u32,
) -> Result<u64, AcceptError> {
table.register(addr, backlog)
}
#[derive(Debug)]
pub struct StdTcpAcceptor {
listener: std::net::TcpListener,
local_addr: NetAddr,
next_conn_id: u64,
}
impl StdTcpAcceptor {
pub fn bind(addr: NetAddr) -> Result<Self, AcceptError> {
let sock_addr = netaddr_to_socketaddr(&addr)
.ok_or_else(|| AcceptError::BindFailed("unsupported address family".to_string()))?;
let listener = std::net::TcpListener::bind(sock_addr).map_err(|e| match e.kind() {
std::io::ErrorKind::AddrInUse => AcceptError::AddrInUse,
std::io::ErrorKind::AddrNotAvailable => AcceptError::AddrNotAvailable,
_ => AcceptError::BindFailed(e.to_string()),
})?;
listener.set_nonblocking(true).map_err(|e| {
AcceptError::BindFailed(format!("set nonblocking failed: {}", e))
})?;
let local_sock = listener.local_addr().map_err(|e| {
AcceptError::BindFailed(format!("get local addr failed: {}", e))
})?;
let local_addr = socketaddr_to_netaddr(&local_sock);
Ok(Self {
listener,
local_addr,
next_conn_id: 1,
})
}
#[inline]
pub fn local_addr(&self) -> NetAddr {
self.local_addr
}
#[inline]
pub fn inner(&self) -> &std::net::TcpListener {
&self.listener
}
}
impl Accept for StdTcpAcceptor {
type Conn = std::net::TcpStream;
type Error = AcceptError;
fn try_accept(&mut self) -> Result<Option<AcceptedConn<Self::Conn>>, Self::Error> {
match self.listener.accept() {
Ok((stream, remote_sock)) => {
stream.set_nodelay(true).map_err(|e| {
AcceptError::AcceptFailed(format!("set TCP_NODELAY failed: {}", e))
})?;
let conn_id = self.next_conn_id;
self.next_conn_id = self
.next_conn_id
.checked_add(1)
.ok_or(AcceptError::AcceptFailed(
"conn_id overflow".to_string(),
))?;
let remote = socketaddr_to_netaddr(&remote_sock);
let local = self.local_addr;
Ok(Some(AcceptedConn::new(stream, remote, local, conn_id)))
}
Err(e) => {
if e.kind() == std::io::ErrorKind::WouldBlock {
Ok(None)
} else {
Err(AcceptError::AcceptFailed(e.to_string()))
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::{Read, Write};
use std::net::TcpStream;
use std::thread;
use std::time::Duration;
#[test]
fn test_bind_table_register_deregister() {
let mut table = BindTable::new();
let addr = NetAddr::new_ipv4([127, 0, 0, 1], 8080);
assert!(table.is_empty());
assert_eq!(table.len(), 0);
let id = table.register(addr, 128).expect("register should succeed");
assert_eq!(id, 1);
assert_eq!(table.len(), 1);
assert!(!table.is_empty());
assert!(table.is_listening(addr));
let ep = table.get_by_id(id).expect("endpoint should exist");
assert_eq!(ep.listen_id, id);
assert_eq!(ep.addr, addr);
assert_eq!(ep.backlog, 128);
assert_eq!(ep.active_syn_count, 0);
table.deregister(id);
assert_eq!(table.len(), 0);
assert!(table.is_empty());
assert!(!table.is_listening(addr));
assert!(table.get_by_id(id).is_none());
}
#[test]
fn test_bind_table_addr_in_use() {
let mut table = BindTable::new();
let addr = NetAddr::new_ipv4([10, 0, 0, 1], 443);
table.register(addr, 64).expect("first register");
let result = table.register(addr, 64);
assert!(matches!(result, Err(AcceptError::AddrInUse)));
}
#[test]
fn test_bind_table_full_overflow() {
let mut table = BindTable::with_capacity(2);
let addr1 = NetAddr::new_ipv4([127, 0, 0, 1], 8001);
let addr2 = NetAddr::new_ipv4([127, 0, 0, 1], 8002);
let addr3 = NetAddr::new_ipv4([127, 0, 0, 1], 8003);
assert!(table.register(addr1, 10).is_ok());
assert!(table.register(addr2, 10).is_ok());
let result = table.register(addr3, 10);
assert!(matches!(result, Err(AcceptError::ListenTableFull)));
assert_eq!(table.len(), 2);
}
#[test]
fn test_bind_table_free_list_reuse() {
let mut table = BindTable::with_capacity(4);
let addr1 = NetAddr::new_ipv4([127, 0, 0, 1], 9001);
let addr2 = NetAddr::new_ipv4([127, 0, 0, 1], 9002);
let addr3 = NetAddr::new_ipv4([127, 0, 0, 1], 9003);
let id1 = table.register(addr1, 10).expect("ok");
let id2 = table.register(addr2, 10).expect("ok");
assert_eq!(id1, 1);
assert_eq!(id2, 2);
table.deregister(id1);
assert_eq!(table.free_list.len(), 1);
let id3 = table.register(addr3, 10).expect("ok");
assert_eq!(id3, 3);
assert_eq!(table.len(), 2);
assert!(table.is_listening(addr2));
assert!(table.is_listening(addr3));
assert!(!table.is_listening(addr1));
}
#[test]
fn test_bind_table_iter() {
let mut table = BindTable::new();
let addr1 = NetAddr::new_ipv4([127, 0, 0, 1], 7001);
let addr2 = NetAddr::new_ipv4([127, 0, 0, 1], 7002);
table.register(addr1, 32).expect("ok");
table.register(addr2, 32).expect("ok");
let collected: Vec<_> = table.iter().map(|e| e.addr).collect();
assert_eq!(collected.len(), 2);
assert!(collected.contains(&addr1));
assert!(collected.contains(&addr2));
}
#[test]
fn test_bind_table_deregister_nonexistent_idempotent() {
let mut table = BindTable::new();
table.deregister(999);
assert!(table.is_empty());
}
#[test]
fn test_netaddr_v4_to_socketaddr() {
let net = NetAddr::new_ipv4([192, 168, 1, 100], 8080);
let sock = netaddr_to_socketaddr(&net).expect("should convert");
match sock {
std::net::SocketAddr::V4(v4) => {
assert_eq!(v4.ip().octets(), [192, 168, 1, 100]);
assert_eq!(v4.port(), 8080);
}
_ => panic!("expected SocketAddrV4"),
}
}
#[test]
fn test_netaddr_v6_to_socketaddr() {
let ip_bytes: [u8; 16] = [
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
0x00, 0x01,
];
let net = NetAddr::new_ipv6(ip_bytes, 443);
let sock = netaddr_to_socketaddr(&net).expect("should convert");
match sock {
std::net::SocketAddr::V6(v6) => {
assert_eq!(v6.ip().octets(), ip_bytes);
assert_eq!(v6.port(), 443);
}
_ => panic!("expected SocketAddrV6"),
}
}
#[test]
fn test_socketaddr_to_netaddr_roundtrip() {
let original = NetAddr::new_ipv4([10, 20, 30, 40], 12345);
let sock = netaddr_to_socketaddr(&original).expect("convert");
let back = socketaddr_to_netaddr(&sock);
assert_eq!(original, back);
}
#[test]
fn test_accept_error_display_variants() {
let cases = vec![
(AcceptError::BindFailed("test".into()), "bind failed: test"),
(
AcceptError::AcceptFailed("eagain".into()),
"accept failed: eagain",
),
(
AcceptError::ListenTableFull,
"listen table full: capacity=64",
),
(AcceptError::EndpointNotFound, "endpoint not found in bind table"),
(AcceptError::AddrInUse, "address already in use"),
(AcceptError::AddrNotAvailable, "address not available"),
(AcceptError::Io("io msg".into()), "io error: io msg"),
];
for (err, expected_substr) in cases {
let msg = format!("{}", err);
assert!(
msg.contains(expected_substr),
"error '{}' should contain '{}'",
msg,
expected_substr
);
}
}
#[test]
fn test_accept_error_clone_eq() {
let e1 = AcceptError::BindFailed("x".into());
let e2 = e1.clone();
assert_eq!(e1, e2);
let e3 = AcceptError::ListenTableFull;
let e4 = AcceptError::ListenTableFull;
assert_eq!(e3, e4);
assert_ne!(e1, e3);
}
#[test]
fn test_io_error_conversion() {
let addr_in_use = std::io::Error::new(std::io::ErrorKind::AddrInUse, "in use");
let ae: AcceptError = addr_in_use.into();
assert!(matches!(ae, AcceptError::AddrInUse));
let addr_not_avail =
std::io::Error::new(std::io::ErrorKind::AddrNotAvailable, "not avail");
let ae2: AcceptError = addr_not_avail.into();
assert!(matches!(ae2, AcceptError::AddrNotAvailable));
let other = std::io::Error::new(std::io::ErrorKind::PermissionDenied, "denied");
let ae3: AcceptError = other.into();
assert!(matches!(ae3, AcceptError::Io(_)));
}
#[test]
fn test_std_tcp_acceptor_bind_accept() {
let bind_addr = NetAddr::new_ipv4([127, 0, 0, 1], 0);
let mut acceptor = StdTcpAcceptor::bind(bind_addr).expect("bind should succeed");
let local = acceptor.local_addr();
assert_eq!(local.family(), AddressFamily::Ipv4);
assert_eq!(local.ipv4_bytes(), [127, 0, 0, 1]);
assert_ne!(local.port(), 0);
let result = acceptor.try_accept().expect("try_accept should not error");
assert!(result.is_none(), "should be none with no connection");
let port = local.port();
let handle = thread::spawn(move || {
thread::sleep(Duration::from_millis(50));
let sock_addr = std::net::SocketAddrV4::new(
std::net::Ipv4Addr::new(127, 0, 0, 1),
port,
);
let mut stream = TcpStream::connect(sock_addr).expect("connect ok");
stream
.set_read_timeout(Some(Duration::from_millis(200)))
.expect("set read timeout");
stream.write_all(b"hello zenith").expect("write ok");
let _ = stream.shutdown(std::net::Shutdown::Write);
let mut buf = [0u8; 16];
let _ = stream.read(&mut buf);
});
let mut accepted = loop {
match acceptor.try_accept() {
Ok(Some(a)) => break a,
Ok(None) => thread::sleep(Duration::from_millis(10)),
Err(e) => panic!("accept error: {}", e),
}
};
assert_eq!(accepted.conn_id, 1);
assert_eq!(accepted.local, local);
assert_eq!(accepted.remote.family(), AddressFamily::Ipv4);
assert_eq!(accepted.remote.ipv4_bytes(), [127, 0, 0, 1]);
let nodelay = accepted
.conn
.nodelay()
.expect("get nodelay");
assert!(nodelay, "TCP_NODELAY should be set");
accepted
.conn
.set_read_timeout(Some(Duration::from_millis(500)))
.expect("set read timeout");
let mut buf = [0u8; 32];
let n = accepted
.conn
.read(&mut buf)
.expect("read ok");
assert_eq!(&buf[..n], b"hello zenith");
handle.join().expect("thread ok");
}
#[test]
fn test_std_tcp_acceptor_bind_addr_in_use() {
let bind_addr = NetAddr::new_ipv4([127, 0, 0, 1], 0);
let acc1 = StdTcpAcceptor::bind(bind_addr).expect("first bind ok");
let port = acc1.local_addr().port();
let addr2 = NetAddr::new_ipv4([127, 0, 0, 1], port);
let result = StdTcpAcceptor::bind(addr2);
assert!(
matches!(result, Err(AcceptError::AddrInUse)),
"expected AddrInUse, got {:?}",
result.err()
);
}
#[test]
fn test_accepted_conn_into_parts() {
let conn = ();
let remote = NetAddr::new_ipv4([1, 2, 3, 4], 1000);
let local = NetAddr::new_ipv4([5, 6, 7, 8], 2000);
let ac = AcceptedConn::new(conn, remote, local, 42);
let (_, r, l, id) = ac.into_parts();
assert_eq!(r, remote);
assert_eq!(l, local);
assert_eq!(id, 42);
}
#[test]
fn test_register_tcp_listener_hook() {
let mut table = BindTable::new();
let addr = NetAddr::new_ipv4([0, 0, 0, 0], 80);
let id = register_tcp_listener(&mut table, addr, 128).expect("register ok");
assert_eq!(id, 1);
assert!(table.is_listening(addr));
let ep = table.get_by_id(id).expect("get ok");
assert_eq!(ep.backlog, 128);
}
#[test]
fn test_listen_endpoint_derives() {
let addr = NetAddr::new_ipv4([10, 0, 0, 1], 80);
let ep = ListenEndpoint {
listen_id: 1,
addr,
backlog: 64,
active_syn_count: 0,
};
let ep2 = ep;
assert_eq!(ep, ep2);
let _ = format!("{:?}", ep);
}
#[test]
fn test_accepted_conn_debug() {
let ac = AcceptedConn::new(
42i32,
NetAddr::new_ipv4([1, 2, 3, 4], 1),
NetAddr::new_ipv4([5, 6, 7, 8], 2),
7,
);
let debug = format!("{:?}", ac);
assert!(debug.contains("conn_id: 7"));
}
#[test]
fn test_bind_table_get_by_id_mut() {
let mut table = BindTable::new();
let addr = NetAddr::new_ipv4([127, 0, 0, 1], 9090);
let id = table.register(addr, 10).expect("ok");
{
let ep = table.get_by_id_mut(id).expect("mut borrow");
ep.active_syn_count = 5;
}
let ep = table.get_by_id(id).expect("read back");
assert_eq!(ep.active_syn_count, 5);
}
#[test]
fn test_std_tcp_acceptor_consecutive_conn_ids() {
let bind_addr = NetAddr::new_ipv4([127, 0, 0, 1], 0);
let mut acceptor = StdTcpAcceptor::bind(bind_addr).expect("bind ok");
let port = acceptor.local_addr().port();
let spawn_client = |p: u16, delay_ms: u64| {
thread::spawn(move || {
thread::sleep(Duration::from_millis(delay_ms));
let sock = std::net::SocketAddrV4::new(
std::net::Ipv4Addr::new(127, 0, 0, 1),
p,
);
let _s = TcpStream::connect(sock).expect("connect");
thread::sleep(Duration::from_millis(100));
})
};
let h1 = spawn_client(port, 10);
let acc1 = loop {
if let Some(a) = acceptor.try_accept().expect("accept") {
break a;
}
thread::sleep(Duration::from_millis(5));
};
assert_eq!(acc1.conn_id, 1);
let h2 = spawn_client(port, 10);
let acc2 = loop {
if let Some(a) = acceptor.try_accept().expect("accept") {
break a;
}
thread::sleep(Duration::from_millis(5));
};
assert_eq!(acc2.conn_id, 2);
h1.join().expect("h1 ok");
h2.join().expect("h2 ok");
}
#[test]
fn test_bind_table_with_capacity() {
let t = BindTable::with_capacity(16);
assert_eq!(t.capacity(), 16);
assert!(t.is_empty());
}
#[test]
fn test_bind_table_default() {
let t: BindTable = Default::default();
assert_eq!(t.capacity(), BIND_TABLE_CAPACITY);
}
}