diff --git a/crates/server/src/lib.rs b/crates/server/src/lib.rs index e5ae0bce..b36c04ab 100755 --- a/crates/server/src/lib.rs +++ b/crates/server/src/lib.rs @@ -248,7 +248,8 @@ pub async fn start_server( // connection. If it's plain HTTP (not 0x16 TLS ClientHello), write // a 301 redirect to HTTPS and reject the connection before the TLS // handshake. This gives users a helpful redirect instead of a - // confusing TLS handshake failure. + // confusing TLS handshake failure. The same acceptor sets + // TCP_NODELAY on every accepted socket before the handshake. let redirect_acceptor = HttpsRedirectAcceptor { addr: local_addr }; axum_server::from_tcp_rustls(std_listener, rustls_config)? .map(|tls| tls.acceptor(redirect_acceptor)) @@ -258,7 +259,7 @@ pub async fn start_server( } else { tracing::info!("extenddb listening on {local_addr}"); axum::serve( - listener, + nodelay_listener(listener), app.into_make_service_with_connect_info::(), ) .with_graceful_shutdown(graceful_shutdown(pid_file.clone())) @@ -298,12 +299,51 @@ fn cleanup_pid_file(pid_file: Option<&std::path::Path>) { } } +/// Set `TCP_NODELAY` on an accepted socket. +/// +/// With Nagle's algorithm on, the kernel holds a small write while an earlier +/// small write on the connection is unacknowledged, and the client's delayed +/// ACK leaves with its next request on that connection or when its timer +/// fires (about 40 ms on Linux). Once one response is written while the +/// previous one is unacknowledged, each response waits for the client's next +/// request, and the reply to that request is held behind it in turn. The +/// response latency then equals the gap between the client's requests on the +/// connection, for every gap under the delayed-ACK timer. `TCP_NODELAY` sends +/// each write as it is made. The option lives on the TCP socket, so under TLS +/// it is set before the handshake and applies below the TLS stream. A failure +/// to set it is logged at debug and the connection is served with the kernel +/// default. +fn set_tcp_nodelay(stream: &tokio::net::TcpStream) { + if let Err(e) = stream.set_nodelay(true) { + tracing::debug!("Failed to set TCP_NODELAY on accepted connection: {e}"); + } +} + +/// The plaintext listener: a `TcpListener` whose accepted sockets carry +/// `TCP_NODELAY`. +type NodelayListener = axum::serve::TapIo; + +/// Wrap the plaintext listener so every accepted socket carries `TCP_NODELAY`. +/// +/// `axum::serve` does not set the option itself; `ListenerExt::tap_io` is its +/// hook for per-connection socket options. +fn nodelay_listener(listener: tokio::net::TcpListener) -> NodelayListener { + use axum::serve::ListenerExt; + listener.tap_io(|stream| set_tcp_nodelay(stream)) +} + /// AI-2: Acceptor that detects plain HTTP connections on the TLS port. /// -/// Peeks the first byte of each connection. If it's `0x16` (TLS `ClientHello`), -/// the connection passes through to the TLS acceptor. If it's a plain HTTP -/// verb, a 301 redirect to `https://` is written and the connection is -/// rejected with an IO error (which `axum_server` handles by dropping it). +/// Sets `TCP_NODELAY` on the accepted socket first (see [`set_tcp_nodelay`]). +/// Then peeks the first byte of the connection. If it's `0x16` (TLS +/// `ClientHello`), the connection passes through to the TLS acceptor. If it's +/// a plain HTTP verb, a 301 redirect to `https://` is written and the +/// connection is rejected with an IO error (which `axum_server` handles by +/// dropping it). +/// +/// `axum_server` 0.8 ships a `NoDelayAcceptor`, but it cannot be stacked under +/// this acceptor without a generic wrapper, so this acceptor sets the option +/// itself, and logs instead of rejecting on failure. #[derive(Clone)] struct HttpsRedirectAcceptor { addr: std::net::SocketAddr, @@ -323,6 +363,7 @@ impl axum_server::accept::Accept fn accept(&self, stream: tokio::net::TcpStream, service: S) -> Self::Future { let addr = self.addr; Box::pin(async move { + set_tcp_nodelay(&stream); let mut peek_buf = [0u8; 1]; match stream.peek(&mut peek_buf).await { Ok(0) => { @@ -393,3 +434,62 @@ async fn shutdown_signal() { let _ = tokio::signal::ctrl_c().await; tracing::info!("Shutdown signal received, draining connections..."); } + +#[cfg(test)] +mod tests { + use super::*; + use tokio::io::AsyncWriteExt; + use tokio::net::{TcpListener, TcpStream}; + + /// Bind a loopback listener and connect one client to it. + async fn bind_and_connect() -> (TcpListener, TcpStream) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("bind loopback"); + let addr = listener.local_addr().expect("local addr"); + let client = TcpStream::connect(addr).await.expect("connect"); + (listener, client) + } + + /// The TLS path: the acceptor that runs before the TLS handshake sets + /// `TCP_NODELAY` on the accepted socket, read back from the socket the + /// acceptor hands on. + #[tokio::test] + async fn tls_acceptor_sets_tcp_nodelay_before_the_handshake() { + use axum_server::accept::Accept; + + let (listener, mut client) = bind_and_connect().await; + // A TLS ClientHello starts with the handshake record type. + client + .write_all(&[0x16]) + .await + .expect("write ClientHello byte"); + let (accepted, _) = listener.accept().await.expect("accept"); + assert!( + !accepted.nodelay().expect("read TCP_NODELAY"), + "a fresh accepted socket has Nagle on" + ); + + let acceptor = HttpsRedirectAcceptor { + addr: listener.local_addr().expect("local addr"), + }; + let (stream, ()) = acceptor + .accept(accepted, ()) + .await + .expect("a TLS ClientHello passes through to the TLS acceptor"); + assert!(stream.nodelay().expect("read TCP_NODELAY")); + } + + /// The plaintext path: the listener handed to `axum::serve` sets + /// `TCP_NODELAY` on every socket it accepts. + #[tokio::test] + async fn plaintext_listener_sets_tcp_nodelay_on_accepted_sockets() { + use axum::serve::Listener; + + let (listener, client) = bind_and_connect().await; + let mut listener = nodelay_listener(listener); + let (accepted, peer) = listener.accept().await; + assert_eq!(peer, client.local_addr().expect("client addr")); + assert!(accepted.nodelay().expect("read TCP_NODELAY")); + } +}