use std::io;
use std::net::SocketAddr;
use std::pin::Pin;
use std::task::{Context, Poll};
use axum::serve::{Listener, ListenerExt, TapIo};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::{mpsc, watch};
use tokio_rustls::server::TlsStream;
use tracing::{debug, warn};
use crate::tls::TlsSettings;
const MAX_PENDING_CONNECTIONS: usize = 256;
const ACCEPT_ERROR_BACKOFF: std::time::Duration = std::time::Duration::from_secs(1);
pub enum MaybeTls {
Plain(TcpStream),
Tls(Box<TlsStream<TcpStream>>),
}
impl AsyncRead for MaybeTls {
fn poll_read(
self: Pin<&mut Self>,
context: &mut Context<'_>,
buffer: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Plain(stream) => Pin::new(stream).poll_read(context, buffer),
Self::Tls(stream) => Pin::new(stream.as_mut()).poll_read(context, buffer),
}
}
}
impl AsyncWrite for MaybeTls {
fn poll_write(
self: Pin<&mut Self>,
context: &mut Context<'_>,
buffer: &[u8],
) -> Poll<io::Result<usize>> {
match self.get_mut() {
Self::Plain(stream) => Pin::new(stream).poll_write(context, buffer),
Self::Tls(stream) => Pin::new(stream.as_mut()).poll_write(context, buffer),
}
}
fn poll_flush(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Plain(stream) => Pin::new(stream).poll_flush(context),
Self::Tls(stream) => Pin::new(stream.as_mut()).poll_flush(context),
}
}
fn poll_shutdown(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.get_mut() {
Self::Plain(stream) => Pin::new(stream).poll_shutdown(context),
Self::Tls(stream) => Pin::new(stream.as_mut()).poll_shutdown(context),
}
}
fn poll_write_vectored(
self: Pin<&mut Self>,
context: &mut Context<'_>,
buffers: &[io::IoSlice<'_>],
) -> Poll<io::Result<usize>> {
match self.get_mut() {
Self::Plain(stream) => Pin::new(stream).poll_write_vectored(context, buffers),
Self::Tls(stream) => Pin::new(stream.as_mut()).poll_write_vectored(context, buffers),
}
}
fn is_write_vectored(&self) -> bool {
match self {
Self::Plain(stream) => stream.is_write_vectored(),
Self::Tls(stream) => stream.is_write_vectored(),
}
}
}
pub enum SocketCommand {
Serve(TcpListener),
Close,
}
pub struct ListenerHandle {
sockets: mpsc::UnboundedSender<SocketCommand>,
tls: watch::Sender<Option<TlsSettings>>,
}
impl ListenerHandle {
pub fn serve(&self, listener: TcpListener) {
let _ = self.sockets.send(SocketCommand::Serve(listener));
}
pub fn close(&self) {
let _ = self.sockets.send(SocketCommand::Close);
}
pub fn set_tls(&self, settings: Option<TlsSettings>) {
self.tls.send_replace(settings);
}
}
pub fn bind_blocking(address: &str) -> io::Result<TcpListener> {
let listener = std::net::TcpListener::bind(address)?;
listener.set_nonblocking(true)?;
TcpListener::from_std(listener)
}
pub struct RoleListener {
incoming: mpsc::Receiver<(MaybeTls, SocketAddr)>,
local_addr: SocketAddr,
bound: watch::Receiver<SocketAddr>,
}
pub type RoleSocket = TapIo<RoleListener, fn(&mut MaybeTls)>;
#[must_use]
pub fn spawn(
role: &'static str,
initial: Option<TcpListener>,
tls: Option<TlsSettings>,
) -> (RoleSocket, ListenerHandle) {
let (sender, incoming) = mpsc::channel(MAX_PENDING_CONNECTIONS);
let (sockets_tx, sockets_rx) = mpsc::unbounded_channel();
let (tls_tx, tls_rx) = watch::channel(tls);
let first_addr = initial
.as_ref()
.and_then(|listener| listener.local_addr().ok())
.unwrap_or_else(|| SocketAddr::from(([0, 0, 0, 0], 0)));
let (bound_tx, bound_rx) = watch::channel(first_addr);
tokio::spawn(async move {
let mut current = initial;
let mut commands = Some(sockets_rx);
loop {
if current.is_none() && commands.is_none() {
sender.closed().await;
debug!(
event = "server_accept_loop_ended",
outcome = "success",
listener = role
);
return;
}
let Ok(permit) = sender.clone().reserve_owned().await else {
debug!(
event = "server_accept_loop_ended",
outcome = "success",
listener = role
);
return;
};
let next = match (current.as_ref(), commands.as_mut()) {
(None, None) => unreachable!("checked at the top of the loop"),
(None, Some(commands)) => Next::Command(commands.recv().await),
(Some(listener), None) => accept_one(listener, role).await,
(Some(listener), Some(commands)) => tokio::select! {
biased;
command = commands.recv() => Next::Command(command),
accepted = accept_one(listener, role) => accepted,
},
};
let (stream, peer) = match next {
Next::Connection(connection) => connection,
Next::Command(Some(command)) => {
apply(&mut current, command, &bound_tx, role);
continue;
}
Next::Command(None) => {
commands = None;
continue;
}
Next::Backoff => {
tokio::time::sleep(ACCEPT_ERROR_BACKOFF).await;
continue;
}
};
let settings = tls_rx.borrow().clone();
match settings {
None => {
permit.send((MaybeTls::Plain(stream), peer));
}
Some(TlsSettings {
acceptor,
handshake_timeout,
}) => {
tokio::spawn(async move {
match tokio::time::timeout(handshake_timeout, acceptor.accept(stream)).await
{
Ok(Ok(tls)) => {
permit.send((MaybeTls::Tls(Box::new(tls)), peer));
}
Ok(Err(error)) => {
debug!(event = "tls_handshake_failed", outcome = "failure", peer = %peer, error = %error);
}
Err(_) => {
debug!(event = "tls_handshake_timeout", outcome = "failure", peer = %peer)
}
}
});
}
}
}
});
let listener = RoleListener {
incoming,
local_addr: first_addr,
bound: bound_rx,
}
.tap_io(noop_tap as fn(&mut MaybeTls));
(
listener,
ListenerHandle {
sockets: sockets_tx,
tls: tls_tx,
},
)
}
enum Next {
Connection((TcpStream, SocketAddr)),
Command(Option<SocketCommand>),
Backoff,
}
async fn accept_one(listener: &TcpListener, role: &'static str) -> Next {
match listener.accept().await {
Ok(connection) => Next::Connection(connection),
Err(error) => {
warn!(
event = "server_accept_failed",
outcome = "failure",
listener = role,
error = %error
);
Next::Backoff
}
}
}
fn apply(
current: &mut Option<TcpListener>,
command: SocketCommand,
bound: &watch::Sender<SocketAddr>,
role: &'static str,
) {
match command {
SocketCommand::Serve(listener) => {
if let Ok(address) = listener.local_addr() {
bound.send_replace(address);
}
*current = Some(listener);
}
SocketCommand::Close => {
*current = None;
debug!(
event = "server_socket_closed",
outcome = "success",
listener = role
);
}
}
}
fn noop_tap(_stream: &mut MaybeTls) {}
impl Listener for RoleListener {
type Io = MaybeTls;
type Addr = SocketAddr;
async fn accept(&mut self) -> (Self::Io, Self::Addr) {
match self.incoming.recv().await {
Some(connection) => {
if self.bound.has_changed().unwrap_or(false) {
self.local_addr = *self.bound.borrow_and_update();
}
connection
}
None => {
tracing::error!(
event = "server_acceptor_stopped",
outcome = "failure",
"the accept task ended: no further connection will be served"
);
std::future::pending().await
}
}
}
fn local_addr(&self) -> io::Result<Self::Addr> {
Ok(self.local_addr)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{ServerConfig, TlsConfig};
use crate::testutil::TempDir;
use axum::extract::ConnectInfo;
use axum::routing::get;
use axum::{Router, serve};
use rustls::pki_types::ServerName;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio_rustls::TlsConnector;
fn tls_config(dir: &TempDir, base_url: &str) -> ServerConfig {
ServerConfig {
bind_address: "127.0.0.1:0".to_string(),
base_url: base_url.to_string(),
tls: TlsConfig {
enabled: true,
cert_path: dir.join("server.pem").display().to_string(),
key_path: dir.join("server.key").display().to_string(),
handshake_timeout_ms: 5_000,
},
..ServerConfig::default()
}
}
fn settings(name: &str, timeout: Duration) -> TlsSettings {
let dir = TempDir::new(name);
let acceptor = crate::tls::from_config(&tls_config(&dir, "https://localhost"))
.unwrap()
.unwrap();
drop(dir);
TlsSettings::new(acceptor, timeout)
}
async fn serve_peer(tls: Option<TlsSettings>) -> (u16, ListenerHandle) {
let tcp = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = tcp.local_addr().unwrap().port();
let (socket, handle) = spawn("test", Some(tcp), tls);
let app = Router::new().route(
"/peer",
get(|ConnectInfo(peer): ConnectInfo<SocketAddr>| async move { peer.to_string() }),
);
tokio::spawn(async move {
serve(
socket,
app.into_make_service_with_connect_info::<SocketAddr>(),
)
.await
.unwrap();
});
(port, handle)
}
async fn get_peer(port: u16) -> (String, SocketAddr) {
let config =
crate::challenge::tls_alpn_01::accept_any_client_config(&[b"http/1.1"]).unwrap();
let stream = TcpStream::connect(("127.0.0.1", port)).await.unwrap();
let client_addr = stream.local_addr().unwrap();
let mut tls = TlsConnector::from(config)
.connect(ServerName::try_from("localhost").unwrap(), stream)
.await
.unwrap();
tls.write_all(b"GET /peer HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = String::new();
tls.read_to_string(&mut response).await.unwrap();
(response, client_addr)
}
async fn get_plain(port: u16) -> String {
let mut stream = TcpStream::connect(("127.0.0.1", port)).await.unwrap();
stream
.write_all(b"GET /peer HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = String::new();
stream.read_to_string(&mut response).await.unwrap();
response
}
async fn peer_certificate(port: u16) -> Vec<u8> {
let config =
crate::challenge::tls_alpn_01::accept_any_client_config(&[b"http/1.1"]).unwrap();
let stream = TcpStream::connect(("127.0.0.1", port)).await.unwrap();
let tls = TlsConnector::from(config)
.connect(ServerName::try_from("localhost").unwrap(), stream)
.await
.unwrap();
tls.get_ref()
.1
.peer_certificates()
.expect("the server presented a certificate")[0]
.to_vec()
}
async fn refused(port: u16) -> bool {
matches!(
tokio::time::timeout(
Duration::from_secs(2),
TcpStream::connect(("127.0.0.1", port)),
)
.await,
Ok(Err(_))
)
}
#[tokio::test]
async fn a_replaced_socket_serves_the_new_port_and_releases_the_old() {
let (first, handle) = serve_peer(None).await;
assert!(get_plain(first).await.starts_with("HTTP/1.1 200 OK"));
let replacement = TcpListener::bind("127.0.0.1:0").await.unwrap();
let second = replacement.local_addr().unwrap().port();
handle.serve(replacement);
let response = get_plain(second).await;
assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
assert!(
refused(first).await,
"the old socket must be released, not merely ignored"
);
}
#[tokio::test]
async fn tls_can_be_switched_on_without_the_socket_moving() {
let (port, handle) = serve_peer(None).await;
assert!(get_plain(port).await.starts_with("HTTP/1.1 200 OK"));
handle.set_tls(Some(settings("listener-flip", Duration::from_secs(5))));
let (response, _) = get_peer(port).await;
assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
handle.set_tls(None);
assert!(get_plain(port).await.starts_with("HTTP/1.1 200 OK"));
}
#[tokio::test]
async fn a_closed_role_refuses_connections_and_can_be_reopened() {
let (port, handle) = serve_peer(None).await;
assert!(get_plain(port).await.starts_with("HTTP/1.1 200 OK"));
handle.close();
tokio::task::yield_now().await;
assert!(refused(port).await, "a closed role must not be listening");
let reopened = TcpListener::bind("127.0.0.1:0").await.unwrap();
let again = reopened.local_addr().unwrap().port();
handle.serve(reopened);
assert!(get_plain(again).await.starts_with("HTTP/1.1 200 OK"));
}
#[tokio::test]
async fn a_swapped_certificate_is_served_to_the_next_connection() {
let (port, handle) =
serve_peer(Some(settings("listener-first", Duration::from_secs(5)))).await;
let before = peer_certificate(port).await;
handle.set_tls(Some(settings("listener-second", Duration::from_secs(5))));
let after = peer_certificate(port).await;
assert_ne!(
before, after,
"the connection after the swap must see the new certificate"
);
}
#[tokio::test]
async fn a_request_is_served_with_the_peer_address_intact() {
let (port, _handle) =
serve_peer(Some(settings("listener-peer", Duration::from_secs(5)))).await;
let (response, client_addr) = get_peer(port).await;
assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
assert!(
response.ends_with(&client_addr.to_string()),
"expected the body to be {client_addr}, got {response}"
);
}
#[tokio::test]
async fn a_cleartext_request_keeps_its_peer_address_too() {
let (port, _handle) = serve_peer(None).await;
let mut stream = TcpStream::connect(("127.0.0.1", port)).await.unwrap();
let client_addr = stream.local_addr().unwrap();
stream
.write_all(b"GET /peer HTTP/1.1\r\nHost: localhost\r\nConnection: close\r\n\r\n")
.await
.unwrap();
let mut response = String::new();
stream.read_to_string(&mut response).await.unwrap();
assert!(response.ends_with(&client_addr.to_string()), "{response}");
}
#[tokio::test]
async fn a_failed_handshake_does_not_stop_the_listener() {
let (port, _handle) =
serve_peer(Some(settings("listener-scan", Duration::from_secs(5)))).await;
let mut plain = TcpStream::connect(("127.0.0.1", port)).await.unwrap();
plain.write_all(b"GET / HTTP/1.1\r\n\r\n").await.unwrap();
let mut ignored = Vec::new();
let _ = plain.read_to_end(&mut ignored).await;
let (response, _) = get_peer(port).await;
assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
}
#[tokio::test]
async fn a_stalled_handshake_times_out_without_blocking_others() {
let (port, _handle) =
serve_peer(Some(settings("listener-stall", Duration::from_millis(300)))).await;
let mut stalled = TcpStream::connect(("127.0.0.1", port)).await.unwrap();
let (response, _) = tokio::time::timeout(Duration::from_secs(5), get_peer(port))
.await
.expect("a stalled handshake must not block the accept loop");
assert!(response.starts_with("HTTP/1.1 200 OK"), "{response}");
let mut buffer = [0u8; 1];
let read = tokio::time::timeout(Duration::from_secs(5), stalled.read(&mut buffer))
.await
.expect("the handshake timeout must close the connection");
assert!(
matches!(read, Ok(0) | Err(_)),
"expected EOF after the handshake timeout, got {read:?}"
);
}
#[tokio::test]
async fn bind_blocking_binds_or_says_why_not() {
let listener = bind_blocking("127.0.0.1:0").expect("an ephemeral port must bind");
let port = listener.local_addr().unwrap().port();
let error = bind_blocking(&format!("127.0.0.1:{port}"))
.expect_err("the port is taken, and saying so is what refuses a reload");
assert_eq!(error.kind(), std::io::ErrorKind::AddrInUse);
assert!(
bind_blocking("not-an-address").is_err(),
"an unparseable address must not panic on the reload path"
);
}
}