Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
184 changes: 181 additions & 3 deletions src/server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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(
Expand All @@ -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) => {
Expand Down Expand Up @@ -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::<ConnectionWatch>().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<TcpStream>,
}

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 {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This method name is quite confusing imho. The method does not wait for the connection to close but until either the connection is closed or the something failed or there is still data plus 100mns. Imho:

  • either we make the function blocking until the connection is effectively closed (matches the function name).
  • either we make this function something wait_for_change that waits until there is a change in the connection (new data, an error...). But then we will likely need to surface something to the user.

If we want to go the blocking function we need to figure out a proper way to do that taking into account fun things like "we want to detect that the oxhttp server actually dropped its clone of TcpStream to close the connection for the client". So, maybe doing a clone of the TcpStream is not the way to go. If you want I can give it a try.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If you want I can give it a try.

Please go ahead, you have a better understanding of the different quirks this change needs to take into account

(I am developing an app in which users write SPARQL queries that can take an very long time to complete due to the size of the database and thus canceling an in-flight request would help.)

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,
Expand All @@ -255,10 +322,21 @@ enum ConnectionState {
fn read_body_and_build_response(
request: RequestBuilder,
reader: BufReader<TcpStream>,
stream: &TcpStream,
on_request: &dyn Fn(&mut Request<Body>) -> Response<Body>,
) -> (Response<Body>, 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()) {
Expand Down Expand Up @@ -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<()> {
Expand Down Expand Up @@ -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::<ConnectionWatch>()
.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::<ConnectionWatch>()
.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::<ConnectionWatch>()
.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(())
}
}
Loading