Skip to content
Merged
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
112 changes: 106 additions & 6 deletions crates/server/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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::<std::net::SocketAddr>(),
)
.with_graceful_shutdown(graceful_shutdown(pid_file.clone()))
Expand Down Expand Up @@ -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<tokio::net::TcpListener, fn(&mut tokio::net::TcpStream)>;

/// 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,
Expand All @@ -323,6 +363,7 @@ impl<S: Send + 'static> axum_server::accept::Accept<tokio::net::TcpStream, S>
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) => {
Expand Down Expand Up @@ -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"));
}
}
Loading