o_sfu/runtime/http_server/
server.rs1use 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
32pub(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
45pub(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
79async 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 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 () = &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
197struct 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;