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..636536a 100644 --- a/src/server.rs +++ b/src/server.rs @@ -6,7 +6,7 @@ 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::thread::{sleep, Builder as ThreadBuilder, JoinHandle}; use std::time::Duration; /// 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,73 @@ 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)); +/// 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 || { +/// 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, Debug)] +pub struct ConnectionWatch { + stream: Arc, +} + +impl ConnectionWatch { + /// 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 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 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]; + 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 to avoid busy loops in the caller. + sleep(Duration::from_millis(100)); + false + } + Err(error) => !matches!( + error.kind(), + ErrorKind::TimedOut | ErrorKind::WouldBlock | ErrorKind::Interrupted + ), + } + } +} + #[derive(Eq, PartialEq, Debug, Copy, Clone)] enum ConnectionState { Close, @@ -255,10 +322,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()) { @@ -366,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<()> { @@ -457,4 +536,103 @@ 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 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() + }) + .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_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 + .extensions() + .get::() + .unwrap() + .clone(); + let closed = watch.wait_closed(); + 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 || { + 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(); + 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(()) + } }