use std::io::{self, Read, Write};
use std::net::{SocketAddr, TcpListener, TcpStream};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
#[cfg(unix)]
use std::os::unix::net::{UnixListener, UnixStream};
use yo_reactor::Reactor;
use yo_resp::engine::{Cmd, ConnId, Sink, Wire, pump};
use crate::poll::Poller;
const READ_CHUNK: usize = 16 * 1024;
const SPIN_TURNS: u32 = 256;
const IDLE_WAIT: Duration = Duration::from_millis(20);
const OWED_WAIT: Duration = Duration::from_millis(1);
const LISTENER: u64 = u64::MAX;
const UNIX_LISTENER: u64 = u64::MAX - 1;
enum Sock {
Tcp(TcpStream),
#[cfg(unix)]
Unix(UnixStream),
}
impl Read for Sock {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Sock::Tcp(s) => s.read(buf),
#[cfg(unix)]
Sock::Unix(s) => s.read(buf),
}
}
}
impl Write for Sock {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
Sock::Tcp(s) => s.write(buf),
#[cfg(unix)]
Sock::Unix(s) => s.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
Sock::Tcp(s) => s.flush(),
#[cfg(unix)]
Sock::Unix(s) => s.flush(),
}
}
}
#[cfg(unix)]
impl std::os::fd::AsRawFd for Sock {
fn as_raw_fd(&self) -> std::os::fd::RawFd {
match self {
Sock::Tcp(s) => s.as_raw_fd(),
Sock::Unix(s) => s.as_raw_fd(),
}
}
}
#[cfg(windows)]
impl std::os::windows::io::AsRawSocket for Sock {
fn as_raw_socket(&self) -> std::os::windows::io::RawSocket {
match self {
Sock::Tcp(s) => s.as_raw_socket(),
}
}
}
enum Door {
Tcp(TcpListener),
#[cfg(unix)]
Unix(UnixListener, PathBuf),
}
impl Door {
fn tcp(&self) -> Option<&TcpListener> {
#[cfg(unix)]
{
match self {
Door::Tcp(l) => Some(l),
Door::Unix(..) => None,
}
}
#[cfg(not(unix))]
{
let Door::Tcp(l) = self;
Some(l)
}
}
fn accept(&self) -> io::Result<Sock> {
match self {
Door::Tcp(l) => {
let (stream, _) = l.accept()?;
stream.set_nonblocking(true)?;
let _ = stream.set_nodelay(true);
Ok(Sock::Tcp(stream))
}
#[cfg(unix)]
Door::Unix(l, _) => {
let (stream, _) = l.accept()?;
stream.set_nonblocking(true)?;
Ok(Sock::Unix(stream))
}
}
}
fn token(&self) -> u64 {
match self {
Door::Tcp(_) => LISTENER,
#[cfg(unix)]
Door::Unix(..) => UNIX_LISTENER,
}
}
}
#[cfg(unix)]
impl std::os::fd::AsRawFd for Door {
fn as_raw_fd(&self) -> std::os::fd::RawFd {
match self {
Door::Tcp(l) => l.as_raw_fd(),
Door::Unix(l, _) => l.as_raw_fd(),
}
}
}
#[cfg(windows)]
impl std::os::windows::io::AsRawSocket for Door {
fn as_raw_socket(&self) -> std::os::windows::io::RawSocket {
match self {
Door::Tcp(l) => l.as_raw_socket(),
}
}
}
impl Drop for Door {
fn drop(&mut self) {
#[cfg(unix)]
if let Door::Unix(_, path) = self {
let _ = std::fs::remove_file(path);
}
}
}
#[derive(Default)]
struct Net {
streams: Vec<Option<Sock>>,
dead: Vec<ConnId>,
gone: Vec<ConnId>,
}
impl Net {
fn attach(&mut self, conn: ConnId, stream: Sock) {
if self.streams.len() <= conn as usize {
self.streams.resize_with(conn as usize + 1, || None);
}
self.streams[conn as usize] = Some(stream);
}
fn is_open(&self, conn: ConnId) -> bool {
self.streams.get(conn as usize).is_some_and(Option::is_some)
}
fn read(&mut self, conn: ConnId, buf: &mut [u8]) -> Option<usize> {
let stream = self.streams.get_mut(conn as usize)?.as_mut()?;
match stream.read(buf) {
Ok(0) => None,
Ok(n) => Some(n),
Err(e) if e.kind() == io::ErrorKind::WouldBlock => Some(0),
Err(e) if e.kind() == io::ErrorKind::Interrupted => Some(0),
Err(_) => None,
}
}
}
impl Sink for Net {
fn write(&mut self, conn: ConnId, bytes: &[u8]) -> usize {
let Some(stream) = self.streams.get_mut(conn as usize).and_then(Option::as_mut) else {
return bytes.len();
};
match stream.write(bytes) {
Ok(n) => n,
Err(e) if e.kind() == io::ErrorKind::WouldBlock => 0,
Err(e) if e.kind() == io::ErrorKind::Interrupted => 0,
Err(_) => {
self.dead.push(conn);
bytes.len()
}
}
}
fn closed(&mut self, conn: ConnId) {
if let Some(slot) = self.streams.get_mut(conn as usize) {
*slot = None;
}
self.gone.push(conn);
}
}
pub struct Server {
doors: Vec<Door>,
reactor: Reactor<Wire<Net>>,
poller: Poller,
batch: Vec<Cmd>,
ready: Vec<u64>,
buf: Vec<u8>,
}
impl Server {
pub fn open(addr: Option<SocketAddr>, path: Option<PathBuf>) -> io::Result<Server> {
let mut doors = Vec::new();
if let Some(addr) = addr {
let listener = TcpListener::bind(addr)?;
listener.set_nonblocking(true)?;
doors.push(Door::Tcp(listener));
}
if let Some(path) = path {
doors.push(unix_door(&path)?);
}
if doors.is_empty() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"nothing to listen on: give a port, a socket file, or both",
));
}
let mut poller = Poller::new()?;
for door in &doors {
poller.add(door, door.token())?;
}
Ok(Server {
doors,
reactor: Reactor::inline(Wire::new(Net::default())),
poller,
batch: Vec::with_capacity(64),
ready: Vec::with_capacity(64),
buf: vec![0; READ_CHUNK],
})
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
match self.doors.iter().find_map(Door::tcp) {
Some(l) => l.local_addr(),
None => Err(io::Error::new(
io::ErrorKind::NotFound,
"this server has no port, only a socket file",
)),
}
}
pub fn run(&mut self, stop: &AtomicBool) -> io::Result<()> {
let mut idle = 0u32;
while !stop.load(Ordering::Relaxed) {
let wait = if idle <= SPIN_TURNS {
Duration::ZERO
} else if self.reactor.engine().owed() > 0 {
OWED_WAIT
} else {
IDLE_WAIT
};
self.poller.wait(&mut self.ready, wait)?;
let mut worked = false;
for at in 0..self.ready.len() {
match self.ready[at] {
LISTENER => self.accept_ready(LISTENER)?,
UNIX_LISTENER => self.accept_ready(UNIX_LISTENER)?,
token => self.read_conn(token as ConnId),
}
worked = true;
}
if pump(&mut self.reactor, &mut self.batch) > 0 {
worked = true;
}
self.bury_dead();
self.forget_closed();
if worked {
idle = 0;
} else {
idle = idle.saturating_add(1);
}
}
Ok(())
}
fn accept_ready(&mut self, token: u64) -> io::Result<()> {
let Some(at) = self.doors.iter().position(|d| d.token() == token) else {
return Ok(());
};
loop {
match self.doors[at].accept() {
Ok(stream) => {
let conn = self.reactor.engine_mut().accept();
self.poller.add(&stream, u64::from(conn))?;
self.reactor.engine_mut().sink_mut().attach(conn, stream);
}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => return Ok(()),
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
}
fn read_conn(&mut self, conn: ConnId) {
if !self.reactor.engine().sink().is_open(conn) {
return;
}
loop {
let read = self
.reactor
.engine_mut()
.sink_mut()
.read(conn, &mut self.buf);
match read {
Some(0) => break,
Some(n) => {
self.reactor.engine_mut().feed(conn, &self.buf[..n]);
if n < self.buf.len() {
break;
}
}
None => {
self.reactor.engine_mut().hangup(conn);
break;
}
}
}
}
fn bury_dead(&mut self) {
while let Some(conn) = self.reactor.engine_mut().sink_mut().dead.pop() {
self.reactor.engine_mut().hangup(conn);
}
}
fn forget_closed(&mut self) {
while let Some(conn) = self.reactor.engine_mut().sink_mut().gone.pop() {
if !self.reactor.engine().sink().is_open(conn) {
self.poller.remove(u64::from(conn));
}
}
}
}
#[cfg(unix)]
fn unix_door(path: &Path) -> io::Result<Door> {
let listener = match UnixListener::bind(path) {
Ok(l) => l,
Err(e) if e.kind() == io::ErrorKind::AddrInUse => {
if UnixStream::connect(path).is_ok() {
return Err(io::Error::new(
io::ErrorKind::AddrInUse,
format!(
"{} is a live socket, something is already serving on it",
path.display()
),
));
}
std::fs::remove_file(path)?;
UnixListener::bind(path)?
}
Err(e) => return Err(e),
};
listener.set_nonblocking(true)?;
Ok(Door::Unix(listener, path.to_path_buf()))
}
#[cfg(not(unix))]
fn unix_door(path: &Path) -> io::Result<Door> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
format!(
"{}: this platform has no unix sockets, so serve on a port",
path.display()
),
))
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::AtomicBool;
struct Stopper(Arc<AtomicBool>);
impl Drop for Stopper {
fn drop(&mut self) {
self.0.store(true, Ordering::Relaxed);
}
}
fn served(client: impl FnOnce(SocketAddr) + Send + 'static) {
let mut server = Server::open(
Some("127.0.0.1:0".parse().expect("a literal address")),
None,
)
.expect("a free port");
let addr = server.local_addr().expect("bound");
let stop = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&stop);
let thread = std::thread::spawn(move || {
let _stopper = Stopper(flag);
client(addr);
});
server.run(&stop).expect("the listener stays up");
if let Err(panic) = thread.join() {
std::panic::resume_unwind(panic);
}
}
fn connect(addr: SocketAddr) -> TcpStream {
let s = TcpStream::connect(addr).expect("the server is listening");
s.set_read_timeout(Some(Duration::from_secs(10)))
.expect("a timeout the platform accepts");
s
}
fn read_exact(stream: &mut impl Read, want: usize) -> Vec<u8> {
let mut got = vec![0; want];
stream.read_exact(&mut got).expect("the reply arrives");
got
}
#[test]
fn a_client_gets_its_replies_over_a_real_socket() {
served(|addr| {
let mut client = connect(addr);
client.write_all(b"*1\r\n$4\r\nPING\r\n").expect("sent");
assert_eq!(read_exact(&mut client, 7), b"+PONG\r\n");
client
.write_all(b"*3\r\n$3\r\nSET\r\n$1\r\nk\r\n$5\r\nvalue\r\n")
.expect("sent");
assert_eq!(read_exact(&mut client, 5), b"+OK\r\n");
client
.write_all(b"*2\r\n$3\r\nGET\r\n$1\r\nk\r\n")
.expect("sent");
assert_eq!(read_exact(&mut client, 11), b"$5\r\nvalue\r\n");
});
}
#[test]
fn a_pipeline_comes_back_in_one_piece_and_in_order() {
served(|addr| {
let mut client = connect(addr);
let mut sent = Vec::new();
for _ in 0..64 {
sent.extend_from_slice(b"*2\r\n$4\r\nINCR\r\n$1\r\nn\r\n");
}
client.write_all(&sent).expect("sent");
let mut want = Vec::new();
for i in 1..=64 {
want.extend_from_slice(format!(":{i}\r\n").as_bytes());
}
assert_eq!(read_exact(&mut client, want.len()), want);
});
}
#[test]
fn two_clients_have_their_own_database_and_share_the_store() {
served(|addr| {
let mut a = connect(addr);
let mut b = connect(addr);
a.write_all(b"*2\r\n$6\r\nSELECT\r\n$1\r\n3\r\n")
.expect("sent");
assert_eq!(read_exact(&mut a, 5), b"+OK\r\n");
a.write_all(b"*3\r\n$3\r\nSET\r\n$1\r\nk\r\n$1\r\na\r\n")
.expect("sent");
assert_eq!(read_exact(&mut a, 5), b"+OK\r\n");
b.write_all(b"*2\r\n$3\r\nGET\r\n$1\r\nk\r\n")
.expect("sent");
assert_eq!(read_exact(&mut b, 5), b"$-1\r\n");
b.write_all(b"*3\r\n$3\r\nSET\r\n$1\r\nk\r\n$1\r\nb\r\n")
.expect("sent");
assert_eq!(read_exact(&mut b, 5), b"+OK\r\n");
a.write_all(b"*2\r\n$3\r\nGET\r\n$1\r\nk\r\n")
.expect("sent");
assert_eq!(read_exact(&mut a, 7), b"$1\r\na\r\n");
});
}
#[test]
fn quit_is_answered_and_then_the_socket_closes() {
served(|addr| {
let mut client = connect(addr);
client.write_all(b"*1\r\n$4\r\nQUIT\r\n").expect("sent");
let mut rest = Vec::new();
client
.read_to_end(&mut rest)
.expect("the server closes rather than leaving it open");
assert_eq!(rest, b"+OK\r\n");
});
}
#[test]
fn a_command_arriving_in_two_packets_is_one_command() {
served(|addr| {
let mut client = connect(addr);
client
.write_all(b"*3\r\n$3\r\nSET\r\n$3\r\nkey\r\n$5\r\nv")
.expect("sent");
std::thread::sleep(Duration::from_millis(20));
client.write_all(b"alue\r\n").expect("sent");
assert_eq!(read_exact(&mut client, 5), b"+OK\r\n");
client
.write_all(b"*2\r\n$3\r\nGET\r\n$3\r\nkey\r\n")
.expect("sent");
assert_eq!(read_exact(&mut client, 11), b"$5\r\nvalue\r\n");
});
}
#[test]
fn a_client_that_drops_frees_its_slot() {
served(|addr| {
for _ in 0..8 {
let mut client = connect(addr);
client
.write_all(b"*2\r\n$4\r\nINCR\r\n$1\r\nn\r\n")
.expect("sent");
let mut reply = [0; 16];
let n = client.read(&mut reply).expect("a reply");
assert!(reply[..n].starts_with(b":"), "{:?}", &reply[..n]);
}
let mut last = connect(addr);
last.write_all(b"*2\r\n$3\r\nGET\r\n$1\r\nn\r\n")
.expect("sent");
assert_eq!(read_exact(&mut last, 7), b"$1\r\n8\r\n");
});
}
#[cfg(unix)]
mod unix {
use super::*;
use std::os::unix::net::UnixStream;
fn socket_path(name: &str) -> PathBuf {
let mut p = std::env::temp_dir();
p.push(format!("yodb-test-{name}-{}.sock", std::process::id()));
let _ = std::fs::remove_file(&p);
p
}
fn served_unix(name: &str, client: impl FnOnce(PathBuf) + Send + 'static) {
let path = socket_path(name);
let mut server = Server::open(None, Some(path.clone())).expect("a fresh path");
let stop = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&stop);
let theirs = path.clone();
let thread = std::thread::spawn(move || {
let _stopper = Stopper(flag);
client(theirs);
});
server.run(&stop).expect("the listener stays up");
if let Err(panic) = thread.join() {
std::panic::resume_unwind(panic);
}
}
#[test]
fn a_client_gets_its_replies_over_a_socket_file() {
served_unix("replies", |path| {
let mut client = UnixStream::connect(&path).expect("the server is listening");
client
.set_read_timeout(Some(Duration::from_secs(10)))
.expect("a timeout the platform accepts");
client.write_all(b"*1\r\n$4\r\nPING\r\n").expect("sent");
assert_eq!(read_exact(&mut client, 7), b"+PONG\r\n");
client
.write_all(b"*3\r\n$3\r\nSET\r\n$1\r\nk\r\n$5\r\nvalue\r\n")
.expect("sent");
assert_eq!(read_exact(&mut client, 5), b"+OK\r\n");
client
.write_all(b"*2\r\n$3\r\nGET\r\n$1\r\nk\r\n")
.expect("sent");
assert_eq!(read_exact(&mut client, 11), b"$5\r\nvalue\r\n");
});
}
#[test]
fn the_port_and_the_socket_file_are_the_same_server() {
let path = socket_path("both");
let mut server = Server::open(
Some("127.0.0.1:0".parse().expect("a literal address")),
Some(path.clone()),
)
.expect("a free port and a fresh path");
let addr = server.local_addr().expect("bound");
let stop = Arc::new(AtomicBool::new(false));
let flag = Arc::clone(&stop);
let thread = std::thread::spawn(move || {
let _stopper = Stopper(flag);
let mut over_tcp = connect(addr);
over_tcp
.write_all(b"*3\r\n$3\r\nSET\r\n$1\r\nk\r\n$4\r\nboth\r\n")
.expect("sent");
assert_eq!(read_exact(&mut over_tcp, 5), b"+OK\r\n");
let mut over_file = UnixStream::connect(&path).expect("listening there too");
over_file
.set_read_timeout(Some(Duration::from_secs(10)))
.expect("a timeout the platform accepts");
over_file
.write_all(b"*2\r\n$3\r\nGET\r\n$1\r\nk\r\n")
.expect("sent");
assert_eq!(read_exact(&mut over_file, 10), b"$4\r\nboth\r\n");
});
server.run(&stop).expect("the listener stays up");
if let Err(panic) = thread.join() {
std::panic::resume_unwind(panic);
}
}
#[test]
fn a_leftover_socket_file_is_cleared() {
let path = socket_path("leftover");
{
let _first = Server::open(None, Some(path.clone())).expect("a fresh path");
}
drop(std::os::unix::net::UnixListener::bind(&path).expect("bound"));
assert!(path.exists(), "the leftover is there");
let second = Server::open(None, Some(path.clone()));
assert!(second.is_ok(), "{:?}", second.err());
}
#[test]
fn a_live_socket_file_is_not_stolen() {
let path = socket_path("live");
let _first = Server::open(None, Some(path.clone())).expect("a fresh path");
let e = match Server::open(None, Some(path.clone())) {
Ok(_) => panic!("something is already serving there"),
Err(e) => e,
};
assert_eq!(e.kind(), io::ErrorKind::AddrInUse, "{e}");
}
#[test]
fn the_socket_file_is_removed_on_the_way_out() {
let path = socket_path("cleanup");
{
let _server = Server::open(None, Some(path.clone())).expect("a fresh path");
assert!(path.exists(), "it is there while the server is");
}
assert!(!path.exists(), "and gone once the server is dropped");
}
#[test]
fn a_server_with_no_door_is_refused() {
let e = match Server::open(None, None) {
Ok(_) => panic!("nothing to listen on"),
Err(e) => e,
};
assert_eq!(e.kind(), io::ErrorKind::InvalidInput);
}
}
}