Skip to main content

o_sfu/runtime/http_server/
server.rs

1//! HTTP listener admission, header deadlines and graceful shutdown.
2
3use std::{
4    io,
5    net::SocketAddr,
6    pin::{Pin, pin},
7    sync::Arc,
8    task::{Context, Poll},
9    time::Duration,
10};
11
12use axum::{Router, extract::ConnectInfo, serve::Listener};
13use hyper::{body::Incoming, server::conn::http1, service::service_fn};
14use hyper_util::rt::{TokioIo, TokioTimer};
15use tokio::{
16    io::{AsyncRead, AsyncWrite, ReadBuf},
17    net::{TcpListener, TcpStream},
18    sync::{OwnedSemaphorePermit, Semaphore},
19    task::JoinSet,
20    time::{Instant, sleep_until},
21};
22use tokio_util::sync::CancellationToken;
23use tower::ServiceExt;
24use tracing::{debug, info, warn};
25
26use super::app;
27use crate::{
28    config::HttpConfig,
29    runtime::{RuntimeMetrics, RuntimeState, telemetry::schema::event as telemetry_event},
30};
31
32/// Binds the configured HTTP address and serves until shutdown completes.
33///
34/// # Errors
35///
36/// Returns [`io::Error`] if binding or reading the listener address fails.
37pub(crate) async fn serve_http(
38    state: RuntimeState,
39    shutdown_token: CancellationToken,
40) -> io::Result<()> {
41    let listener = TcpListener::bind(state.config.http.bind_address).await?;
42    serve_http_on(listener, state, shutdown_token).await
43}
44
45/// Serves an existing listener with authorization bound to its actual address.
46///
47/// Shutdown stops acceptance and drains HTTP requests. Upgraded WebSocket connections
48/// retain their connection permits and follow the runtime's session shutdown.
49///
50/// # Errors
51///
52/// Returns [`io::Error`] if reading the listener address fails.
53pub(crate) async fn serve_http_on(
54    listener: TcpListener,
55    state: RuntimeState,
56    shutdown_token: CancellationToken,
57) -> io::Result<()> {
58    let local_address = listener.local_addr()?;
59    info!(
60        event = telemetry_event::HTTP_LISTENER_READY,
61        bind_address = %state.config.http.bind_address,
62        local_address = %local_address,
63        trust_proxy_headers = state.config.http.trust_proxy_headers,
64        "booted HTTP and WebSocket listener"
65    );
66    let config = state.config.http.clone();
67    let metrics = Arc::clone(&state.metrics);
68    serve_connections(
69        listener,
70        app(state, local_address),
71        config,
72        metrics,
73        shutdown_token,
74    )
75    .await;
76    Ok(())
77}
78
79/// Owns HTTP connection tasks so cancelling the listener also closes their sockets.
80async fn serve_connections(
81    mut listener: TcpListener,
82    router: Router,
83    config: HttpConfig,
84    metrics: Arc<RuntimeMetrics>,
85    shutdown: CancellationToken,
86) {
87    let permits = Arc::new(Semaphore::new(config.max_http_connections));
88    let mut connections = JoinSet::new();
89    loop {
90        tokio::select! {
91            biased;
92            () = shutdown.cancelled() => break,
93            Some(result) = connections.join_next(), if !connections.is_empty() => {
94                if let Err(error) = result {
95                    warn!(?error, "HTTP connection task failed");
96                }
97            },
98            (stream, remote_address) = Listener::accept(&mut listener) => {
99                let accepted_at = Instant::now();
100                let Ok(permit) = Arc::clone(&permits).try_acquire_owned() else {
101                    metrics.record_http_connection_rejection();
102                    continue;
103                };
104                connections.spawn(serve_connection(
105                    AdmittedSocket { stream, _permit: permit },
106                    router.clone(),
107                    remote_address,
108                    accepted_at + config.header_read_timeout,
109                    config.header_read_timeout,
110                    shutdown.clone(),
111                ));
112            }
113        }
114    }
115    drop(listener);
116    while let Some(result) = connections.join_next().await {
117        if let Err(error) = result {
118            warn!(?error, "HTTP connection task failed during shutdown");
119        }
120    }
121}
122
123#[expect(
124    clippy::significant_drop_tightening,
125    reason = "Hyper must retain the socket permit while polling the connection and transfer it to upgraded IO"
126)]
127async fn serve_connection(
128    socket: AdmittedSocket,
129    router: Router,
130    remote_address: SocketAddr,
131    first_header_deadline: Instant,
132    header_timeout: Duration,
133    shutdown: CancellationToken,
134) {
135    if Instant::now() >= first_header_deadline {
136        return;
137    }
138    let first_headers = CancellationToken::new();
139    let received_headers = first_headers.clone();
140    let service = service_fn(move |mut request: hyper::Request<Incoming>| {
141        // Ready socket IO can outrun timer notifications after a scheduling
142        // stall. First headers must meet the deadline before any router effects.
143        let expired = !received_headers.is_cancelled() && Instant::now() >= first_header_deadline;
144        if !expired {
145            received_headers.cancel();
146        }
147        let router = router.clone();
148        async move {
149            if expired {
150                return Err(io::Error::from(io::ErrorKind::TimedOut));
151            }
152            request.extensions_mut().insert(ConnectInfo(remote_address));
153            router
154                .oneshot(request)
155                .await
156                .map_err(|never| match never {})
157        }
158    });
159    let mut builder = http1::Builder::new();
160    builder
161        .timer(TokioTimer::new())
162        .header_read_timeout(header_timeout);
163    let mut connection = pin!(
164        builder
165            .serve_connection(TokioIo::new(socket), service)
166            .with_upgrades()
167    );
168    let mut deadline = pin!(sleep_until(first_header_deadline));
169    let mut awaiting_headers = true;
170    let mut draining = false;
171    loop {
172        tokio::select! {
173            biased;
174            () = shutdown.cancelled(), if !draining => {
175                if !first_headers.is_cancelled() {
176                    return;
177                }
178                draining = true;
179                connection.as_mut().graceful_shutdown();
180            }
181            () = first_headers.cancelled(), if awaiting_headers => {
182                awaiting_headers = false;
183            }
184            // Hyper starts its header timer when polled. This deadline includes
185            // the time between acceptance and the connection task's first poll.
186            () = &mut deadline, if awaiting_headers => return,
187            result = connection.as_mut() => {
188                if let Err(error) = result {
189                    debug!(?error, "HTTP connection closed");
190                }
191                return;
192            }
193        }
194    }
195}
196
197/// Couples admission to the socket because Hyper transfers the IO into an
198/// upgraded WebSocket before its HTTP connection future completes.
199struct AdmittedSocket {
200    stream: TcpStream,
201    _permit: OwnedSemaphorePermit,
202}
203
204impl AsyncRead for AdmittedSocket {
205    fn poll_read(
206        self: Pin<&mut Self>,
207        context: &mut Context<'_>,
208        buffer: &mut ReadBuf<'_>,
209    ) -> Poll<io::Result<()>> {
210        Pin::new(&mut self.get_mut().stream).poll_read(context, buffer)
211    }
212}
213
214impl AsyncWrite for AdmittedSocket {
215    fn poll_write(
216        self: Pin<&mut Self>,
217        context: &mut Context<'_>,
218        buffer: &[u8],
219    ) -> Poll<io::Result<usize>> {
220        Pin::new(&mut self.get_mut().stream).poll_write(context, buffer)
221    }
222
223    fn poll_flush(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
224        Pin::new(&mut self.get_mut().stream).poll_flush(context)
225    }
226
227    fn poll_shutdown(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<io::Result<()>> {
228        Pin::new(&mut self.get_mut().stream).poll_shutdown(context)
229    }
230
231    fn is_write_vectored(&self) -> bool {
232        self.stream.is_write_vectored()
233    }
234
235    fn poll_write_vectored(
236        self: Pin<&mut Self>,
237        context: &mut Context<'_>,
238        buffers: &[io::IoSlice<'_>],
239    ) -> Poll<io::Result<usize>> {
240        Pin::new(&mut self.get_mut().stream).poll_write_vectored(context, buffers)
241    }
242}
243
244#[cfg(test)]
245#[path = "TESTS/server.rs"]
246mod tests;