From 132ecbbfd358dfbb3fe4156def93a31c123bbc28 Mon Sep 17 00:00:00 2001 From: Jani Simomaa Date: Thu, 17 Sep 2026 13:49:11 +0300 Subject: [PATCH 1/2] expose connection liveness to the handler --- CHANGELOG.md | 7 ++ src/lib.rs | 2 +- src/server.rs | 184 ++++++++++++++++++++++++++++++++++++++++++++++++-- 3 files changed, 188 insertions(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index cc69000..038ed04 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,12 @@ # Changelog +## [Unreleased] + +### Added + +- Server: `ConnectionWatch` in the request extensions lets a handler wait until the client closes the connection. + + ## [0.3.3] - 2026-06-17 ### Changed diff --git a/src/lib.rs b/src/lib.rs index b025ee2..359f16c 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -22,4 +22,4 @@ mod utils; #[cfg(feature = "client")] pub use client::Client; #[cfg(feature = "server")] -pub use server::{ListeningServer, Server}; +pub use server::{ConnectionWatch, ListeningServer, Server}; diff --git a/src/server.rs b/src/server.rs index 0f35edf..070e3b9 100644 --- a/src/server.rs +++ b/src/server.rs @@ -6,8 +6,8 @@ use std::fmt; use std::io::{copy, sink, BufReader, BufWriter, Error, ErrorKind, Result, Write}; use std::net::{SocketAddr, TcpListener, TcpStream}; use std::sync::{Arc, Condvar, Mutex}; -use std::thread::{Builder as ThreadBuilder, JoinHandle}; -use std::time::Duration; +use std::thread::{sleep, Builder as ThreadBuilder, JoinHandle}; +use std::time::{Duration, Instant}; /// An HTTP server. /// @@ -201,7 +201,7 @@ fn accept_request( && expect.as_bytes().eq_ignore_ascii_case(b"100-continue") { stream.write_all(b"HTTP/1.1 100 Continue\r\n\r\n")?; - read_body_and_build_response(request, reader, on_request) + read_body_and_build_response(request, reader, &stream, on_request) } else { ( build_text_response( @@ -215,7 +215,7 @@ fn accept_request( ) } } else { - read_body_and_build_response(request, reader, on_request) + read_body_and_build_response(request, reader, &stream, on_request) } } Err(error) => { @@ -246,6 +246,82 @@ fn accept_request( Ok(()) } +/// Lets a request handler know whether the client connection is still open. +/// +/// The [`Server`] inserts it into the [extensions](Request::extensions) of each request before calling the request handler. +/// A handler that runs a long computation can watch it from another thread and stop the computation when the client is gone. +/// +/// ```no_run +/// use oxhttp::{ConnectionWatch, Server}; +/// use oxhttp::model::{Body, Response}; +/// use std::sync::Arc; +/// use std::sync::atomic::{AtomicBool, Ordering}; +/// use std::thread; +/// +/// Server::new(|request| { +/// let cancelled = Arc::new(AtomicBool::new(false)); +/// if let Some(watch) = request.extensions().get::().cloned() { +/// let cancelled = Arc::clone(&cancelled); +/// thread::spawn(move || { +/// if watch.wait_closed(None) { +/// cancelled.store(true, Ordering::Relaxed); +/// } +/// }); +/// } +/// // ... run a computation that checks `cancelled` ... +/// Response::builder().body(Body::from("done")).unwrap() +/// }); +/// ``` +#[derive(Clone)] +pub struct ConnectionWatch { + stream: Arc, +} + +impl ConnectionWatch { + /// Blocks until the client closes the connection or `timeout` elapses. + /// + /// Returns `true` if the client closed the connection and `false` if the timeout elapsed first. + /// + /// The connection counts as closed when the client shut down its writing side or when the socket reported an error. + /// Bytes the client sends while waiting (the rest of the request body, or a pipelined next request) are left in place. + /// + /// The check wakes up at most every [global timeout](Server::with_global_timeout) of the server, + /// so this method might return up to one global timeout after the given `timeout` when the connection stays open. + pub fn wait_closed(&self, timeout: Option) -> bool { + let deadline = timeout.map(|timeout| Instant::now() + timeout); + let mut buffer = [0; 1]; + loop { + let has_pending_data = match self.stream.peek(&mut buffer) { + Ok(0) => return true, // EOF: the client closed its writing side + Ok(_) => true, + Err(error) => match error.kind() { + ErrorKind::TimedOut | ErrorKind::WouldBlock | ErrorKind::Interrupted => false, + _ => return true, + }, + }; + let remaining = + deadline.map(|deadline| deadline.saturating_duration_since(Instant::now())); + if remaining == Some(Duration::ZERO) { + return false; + } + if has_pending_data { + // The client sent something we must not consume, so peek would return immediately. + // Let's wait a bit before trying again. + let pause = Duration::from_millis(100); + sleep(remaining.map_or(pause, |remaining| remaining.min(pause))); + } + } + } +} + +impl fmt::Debug for ConnectionWatch { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ConnectionWatch") + .field("peer_addr", &self.stream.peer_addr().ok()) + .finish() + } +} + #[derive(Eq, PartialEq, Debug, Copy, Clone)] enum ConnectionState { Close, @@ -255,10 +331,21 @@ enum ConnectionState { fn read_body_and_build_response( request: RequestBuilder, reader: BufReader, + stream: &TcpStream, on_request: &dyn Fn(&mut Request) -> Response, ) -> (Response, ConnectionState) { match decode_request_body(request, reader) { Ok(mut request) => { + match stream.try_clone() { + Ok(stream) => { + request.extensions_mut().insert(ConnectionWatch { + stream: Arc::new(stream), + }); + } + Err(error) => { + eprintln!("OxHTTP TCP error when attempting to clone the stream: {error}"); + } + } let response = on_request(&mut request); // We make sure to finish reading the body if let Err(error) = drain_body(request.body_mut()) { @@ -457,4 +544,93 @@ mod tests { } Ok(()) } + + #[test] + fn test_connection_watch_detects_client_close() -> Result<()> { + let server_port = 9995; + let (sender, receiver) = std::sync::mpsc::channel(); + Server::new(move |request| { + let watch = request + .extensions() + .get::() + .unwrap() + .clone(); + let start = Instant::now(); + let closed = watch.wait_closed(Some(Duration::from_secs(5))); + sender.send((closed, start.elapsed())).unwrap(); + Response::builder().body(Body::from("home")).unwrap() + }) + .bind((Ipv4Addr::LOCALHOST, server_port)) + .with_global_timeout(Duration::from_secs(1)) + .spawn()?; + sleep(Duration::from_millis(100)); // Makes sure the server is up + let mut stream = TcpStream::connect(("127.0.0.1", server_port))?; + stream.write_all(b"GET / HTTP/1.1\nhost: localhost:9995\n\n")?; + sleep(Duration::from_millis(200)); // Makes sure the handler is waiting + drop(stream); + let (closed, elapsed) = receiver.recv().unwrap(); + assert!(closed); + assert!(elapsed < Duration::from_secs(2), "took {elapsed:?}"); + Ok(()) + } + + #[test] + fn test_connection_watch_timeout() -> Result<()> { + let server_port = 9994; + Server::new(|request| { + let watch = request + .extensions() + .get::() + .unwrap() + .clone(); + let closed = watch.wait_closed(Some(Duration::from_millis(100))); + Response::builder() + .body(Body::from(if closed { "closed" } else { "open" })) + .unwrap() + }) + .bind((Ipv4Addr::LOCALHOST, server_port)) + .with_global_timeout(Duration::from_millis(100)) + .spawn()?; + sleep(Duration::from_millis(100)); // Makes sure the server is up + let mut stream = TcpStream::connect(("127.0.0.1", server_port))?; + stream.write_all(b"GET / HTTP/1.1\nhost: localhost:9994\n\n")?; + let expected = b"HTTP/1.1 200 OK\r\ncontent-length: 4\r\n\r\nopen"; + let mut output = vec![b'\0'; expected.len()]; + stream.read_exact(&mut output)?; + assert_eq!(output, expected); + Ok(()) + } + + #[test] + fn test_connection_watch_keeps_pending_bytes() -> Result<()> { + // A watcher thread peeks the socket while the request body is still unread: + // the body must still be readable by the handler. + let server_port = 9993; + Server::new(|request| { + let watch = request + .extensions() + .get::() + .unwrap() + .clone(); + let watcher = + std::thread::spawn(move || watch.wait_closed(Some(Duration::from_millis(300)))); + sleep(Duration::from_millis(200)); + let mut body = String::new(); + request.body_mut().read_to_string(&mut body).unwrap(); + let closed = watcher.join().unwrap(); + assert!(!closed); + Response::builder().body(Body::from(body)).unwrap() + }) + .bind((Ipv4Addr::LOCALHOST, server_port)) + .with_global_timeout(Duration::from_secs(1)) + .spawn()?; + sleep(Duration::from_millis(100)); // Makes sure the server is up + let mut stream = TcpStream::connect(("127.0.0.1", server_port))?; + stream.write_all(b"POST / HTTP/1.1\nhost: localhost:9993\ncontent-length: 4\n\nabcd")?; + let expected = b"HTTP/1.1 200 OK\r\ncontent-length: 4\r\n\r\nabcd"; + let mut output = vec![b'\0'; expected.len()]; + stream.read_exact(&mut output)?; + assert_eq!(output, expected); + Ok(()) + } } From 3cea343ff01222f4309c22f568fec0094d7013a0 Mon Sep 17 00:00:00 2001 From: Jani Simomaa Date: Thu, 24 Sep 2026 14:41:33 +0300 Subject: [PATCH 2/2] ConnectionWatch: bound wait_closed by the global timeout and derive Debug wait_closed no longer takes a timeout parameter. The socket options that a reliable timeout would need are shared with the body reader, so the method now waits at most one global timeout and returns false if the connection is still open. Callers loop with their own stop condition. --- src/server.rs | 84 ++++++++++++++++++++++++++------------------------- 1 file changed, 43 insertions(+), 41 deletions(-) diff --git a/src/server.rs b/src/server.rs index 070e3b9..636536a 100644 --- a/src/server.rs +++ b/src/server.rs @@ -7,7 +7,7 @@ use std::io::{copy, sink, BufReader, BufWriter, Error, ErrorKind, Result, Write} use std::net::{SocketAddr, TcpListener, TcpStream}; use std::sync::{Arc, Condvar, Mutex}; use std::thread::{sleep, Builder as ThreadBuilder, JoinHandle}; -use std::time::{Duration, Instant}; +use std::time::Duration; /// An HTTP server. /// @@ -260,68 +260,59 @@ fn accept_request( /// /// Server::new(|request| { /// let cancelled = Arc::new(AtomicBool::new(false)); +/// let done = Arc::new(AtomicBool::new(false)); /// if let Some(watch) = request.extensions().get::().cloned() { /// let cancelled = Arc::clone(&cancelled); +/// let done = Arc::clone(&done); /// thread::spawn(move || { -/// if watch.wait_closed(None) { -/// cancelled.store(true, Ordering::Relaxed); +/// while !done.load(Ordering::Relaxed) { +/// if watch.wait_closed() { +/// cancelled.store(true, Ordering::Relaxed); +/// return; +/// } /// } /// }); /// } /// // ... run a computation that checks `cancelled` ... +/// done.store(true, Ordering::Relaxed); /// Response::builder().body(Body::from("done")).unwrap() /// }); /// ``` -#[derive(Clone)] +#[derive(Clone, Debug)] pub struct ConnectionWatch { stream: Arc, } impl ConnectionWatch { - /// Blocks until the client closes the connection or `timeout` elapses. + /// Blocks until the client closes the connection or until one wait period ends. /// - /// Returns `true` if the client closed the connection and `false` if the timeout elapsed first. + /// Returns `true` if the client closed the connection and `false` if the wait period ended first. + /// Callers that want to keep watching call this method in a loop and check their own stop condition between calls. /// /// The connection counts as closed when the client shut down its writing side or when the socket reported an error. /// Bytes the client sends while waiting (the rest of the request body, or a pipelined next request) are left in place. /// - /// The check wakes up at most every [global timeout](Server::with_global_timeout) of the server, - /// so this method might return up to one global timeout after the given `timeout` when the connection stays open. - pub fn wait_closed(&self, timeout: Option) -> bool { - let deadline = timeout.map(|timeout| Instant::now() + timeout); + /// The wait period is the [global timeout](Server::with_global_timeout) of the server. + /// It is 100ms if the client already sent bytes that are not read yet. + /// Without a global timeout, this method blocks until the client closes the connection or sends bytes. + pub fn wait_closed(&self) -> bool { let mut buffer = [0; 1]; - loop { - let has_pending_data = match self.stream.peek(&mut buffer) { - Ok(0) => return true, // EOF: the client closed its writing side - Ok(_) => true, - Err(error) => match error.kind() { - ErrorKind::TimedOut | ErrorKind::WouldBlock | ErrorKind::Interrupted => false, - _ => return true, - }, - }; - let remaining = - deadline.map(|deadline| deadline.saturating_duration_since(Instant::now())); - if remaining == Some(Duration::ZERO) { - return false; - } - if has_pending_data { + match self.stream.peek(&mut buffer) { + Ok(0) => true, // EOF: the client closed its writing side + Ok(_) => { // The client sent something we must not consume, so peek would return immediately. - // Let's wait a bit before trying again. - let pause = Duration::from_millis(100); - sleep(remaining.map_or(pause, |remaining| remaining.min(pause))); + // Let's wait a bit to avoid busy loops in the caller. + sleep(Duration::from_millis(100)); + false } + Err(error) => !matches!( + error.kind(), + ErrorKind::TimedOut | ErrorKind::WouldBlock | ErrorKind::Interrupted + ), } } } -impl fmt::Debug for ConnectionWatch { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("ConnectionWatch") - .field("peer_addr", &self.stream.peer_addr().ok()) - .finish() - } -} - #[derive(Eq, PartialEq, Debug, Copy, Clone)] enum ConnectionState { Close, @@ -453,6 +444,7 @@ mod tests { use std::io::Read; use std::net::{Ipv4Addr, Ipv6Addr}; use std::thread::sleep; + use std::time::Instant; #[test] fn test_regular_http_operations() -> Result<()> { @@ -556,7 +548,10 @@ mod tests { .unwrap() .clone(); let start = Instant::now(); - let closed = watch.wait_closed(Some(Duration::from_secs(5))); + let mut closed = false; + while !closed && start.elapsed() < Duration::from_secs(5) { + closed = watch.wait_closed(); + } sender.send((closed, start.elapsed())).unwrap(); Response::builder().body(Body::from("home")).unwrap() }) @@ -575,7 +570,8 @@ mod tests { } #[test] - fn test_connection_watch_timeout() -> Result<()> { + fn test_connection_watch_returns_after_global_timeout() -> Result<()> { + // The client stays silent: wait_closed must return once the global timeout elapses. let server_port = 9994; Server::new(|request| { let watch = request @@ -583,7 +579,7 @@ mod tests { .get::() .unwrap() .clone(); - let closed = watch.wait_closed(Some(Duration::from_millis(100))); + let closed = watch.wait_closed(); Response::builder() .body(Body::from(if closed { "closed" } else { "open" })) .unwrap() @@ -612,8 +608,14 @@ mod tests { .get::() .unwrap() .clone(); - let watcher = - std::thread::spawn(move || watch.wait_closed(Some(Duration::from_millis(300)))); + let watcher = std::thread::spawn(move || { + let start = Instant::now(); + let mut closed = false; + while !closed && start.elapsed() < Duration::from_millis(300) { + closed = watch.wait_closed(); + } + closed + }); sleep(Duration::from_millis(200)); let mut body = String::new(); request.body_mut().read_to_string(&mut body).unwrap();