tachyon_web/server/
http.rs1use crate::http::response::Body;
2#[cfg(feature = "tls")]
3use crate::server::TLS_HANDSHAKE_TIMEOUT;
4use crate::server::{IS_LOCAL_WORKER, REQUEST_TIMEOUT, Server};
5use bytes::Bytes;
6use hyper::body::{Body as HyperBody, Frame, SizeHint};
7use hyper::service::service_fn;
8use hyper::{Request, Response};
9use std::future::Future;
10use std::pin::Pin;
11use std::sync::Arc;
12use std::task::{Context, Poll};
13use tokio::net::TcpListener;
14#[cfg(feature = "tls")]
15use tokio_rustls::TlsAcceptor;
16
17pin_project_lite::pin_project! {
18 struct DeadlineBody {
36 #[pin]
37 inner: hyper::body::Incoming,
38 deadline: Option<Pin<Box<tokio::time::Sleep>>>,
39 }
40}
41
42impl HyperBody for DeadlineBody {
43 type Data = Bytes;
44 type Error = crate::http::error::Error;
45
46 fn poll_frame(
47 self: Pin<&mut Self>,
48 cx: &mut Context<'_>,
49 ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
50 let this = self.project();
51 let deadline = this
52 .deadline
53 .get_or_insert_with(|| Box::pin(tokio::time::sleep(REQUEST_TIMEOUT)));
54 if deadline.as_mut().poll(cx).is_ready() {
55 return Poll::Ready(Some(Err(crate::http::error::Error::Rejection {
56 status: hyper::StatusCode::REQUEST_TIMEOUT,
57 message: "Timed out reading request body".to_string(),
58 })));
59 }
60 match this.inner.poll_frame(cx) {
61 Poll::Ready(Some(Ok(frame))) => Poll::Ready(Some(Ok(frame))),
62 Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(e.into()))),
63 Poll::Ready(None) => Poll::Ready(None),
64 Poll::Pending => Poll::Pending,
65 }
66 }
67
68 fn is_end_stream(&self) -> bool {
69 self.inner.is_end_stream()
70 }
71
72 fn size_hint(&self) -> SizeHint {
73 self.inner.size_hint()
74 }
75}
76
77#[cfg(feature = "http2")]
78#[derive(Clone, Copy, Debug)]
79struct LocalExecutor;
80
81#[cfg(feature = "http2")]
82impl<F> hyper::rt::Executor<F> for LocalExecutor
83where
84 F: Future + Send + 'static,
85 F::Output: Send + 'static,
86{
87 fn execute(&self, fut: F) {
88 IS_LOCAL_WORKER.with(|flag| {
89 if flag.get() {
90 drop(tokio::task::spawn_local(fut));
91 } else {
92 drop(tokio::spawn(fut));
93 }
94 });
95 }
96}
97
98impl<S> Server<S>
99where
100 S: Clone + Send + Sync + 'static,
101{
102 pub async fn serve_http(self, listener: TcpListener) -> Result<(), std::io::Error> {
124 crate::server::enforce_fips_compliance()?;
125 let state = Arc::new(self);
126 let connection_semaphore = Arc::new(tokio::sync::Semaphore::new(state.max_connections));
127
128 #[cfg(all(feature = "http1", feature = "http2"))]
132 let builder = {
133 let mut b = hyper_util::server::conn::auto::Builder::new(LocalExecutor);
136 let _ = b
137 .http1()
138 .timer(hyper_util::rt::TokioTimer::new())
139 .header_read_timeout(REQUEST_TIMEOUT)
140 .keep_alive(true)
141 .max_buf_size(8192)
142 .writev(true);
143 let _ = b
144 .http2()
145 .timer(hyper_util::rt::TokioTimer::new())
146 .initial_stream_window_size(65535)
147 .initial_connection_window_size(1024 * 1024)
148 .max_frame_size(16384)
149 .max_concurrent_streams(200)
150 .keep_alive_timeout(REQUEST_TIMEOUT);
151 #[cfg(feature = "ws")]
153 let _ = b.http2().enable_connect_protocol();
154 b
155 };
156 #[cfg(all(feature = "http1", not(feature = "http2")))]
157 let builder = {
158 let mut b = hyper::server::conn::http1::Builder::new();
160 let _ = b
161 .timer(hyper_util::rt::TokioTimer::new())
162 .header_read_timeout(REQUEST_TIMEOUT)
163 .keep_alive(true)
164 .max_buf_size(8192)
165 .writev(true);
166 b
167 };
168 #[cfg(all(feature = "http2", not(feature = "http1")))]
169 let builder = {
170 let mut b = hyper::server::conn::http2::Builder::new(LocalExecutor);
173 let _ = b
174 .timer(hyper_util::rt::TokioTimer::new())
175 .initial_stream_window_size(65535)
176 .initial_connection_window_size(1024 * 1024)
177 .max_frame_size(16384)
178 .max_concurrent_streams(200)
179 .keep_alive_timeout(REQUEST_TIMEOUT);
180 #[cfg(feature = "ws")]
181 let _ = b.enable_connect_protocol();
182 b
183 };
184
185 loop {
186 let Ok(permit) = connection_semaphore.clone().acquire_owned().await else {
187 break;
188 };
189
190 let (stream, peer) = match listener.accept().await {
191 Ok(c) => c,
192 Err(e) => {
193 drop(permit);
194 tracing::error!("[http] Accept error: {}", e);
195 if crate::server::is_resource_exhaustion(&e) {
196 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
197 }
198 continue;
199 }
200 };
201 let _ = stream.set_nodelay(true);
202 #[cfg(target_os = "linux")]
203 {
204 let sock_ref = socket2::SockRef::from(&stream);
205 let _ = sock_ref.set_tcp_quickack(true);
206 }
207 let state = state.clone();
208 let builder = builder.clone();
209
210 let serve_fut = async move {
211 let io = hyper_util::rt::TokioIo::new(stream);
212 let svc = service_fn(move |req| hyper_handler(state.clone(), req, peer));
213 #[cfg(all(feature = "http1", feature = "http2"))]
214 let result = builder.serve_connection_with_upgrades(io, svc).await;
215 #[cfg(all(feature = "http1", not(feature = "http2")))]
216 let result = builder.serve_connection(io, svc).with_upgrades().await;
217 #[cfg(all(feature = "http2", not(feature = "http1")))]
218 let result = builder.serve_connection(io, svc).await;
219 if let Err(e) = result {
220 tracing::debug!("[http] Connection error: {}", e);
221 }
222 drop(permit);
223 };
224
225 IS_LOCAL_WORKER.with(|flag| {
226 if flag.get() {
227 drop(tokio::task::spawn_local(serve_fut));
228 } else {
229 drop(tokio::spawn(serve_fut));
230 }
231 });
232 }
233 Ok(())
234 }
235
236 #[cfg(feature = "tls")]
244 pub async fn serve_https(
245 self,
246 listener: TcpListener,
247 acceptor: TlsAcceptor,
248 ) -> Result<(), std::io::Error> {
249 crate::server::enforce_fips_compliance()?;
250 let state = Arc::new(self);
251 let connection_semaphore = Arc::new(tokio::sync::Semaphore::new(state.max_connections));
252
253 loop {
254 let Ok(permit) = connection_semaphore.clone().acquire_owned().await else {
255 break;
256 };
257
258 let (tcp_stream, peer) = match listener.accept().await {
259 Ok(c) => c,
260 Err(e) => {
261 drop(permit);
262 tracing::error!("[https] Accept error: {}", e);
263 if crate::server::is_resource_exhaustion(&e) {
264 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
265 }
266 continue;
267 }
268 };
269 let _ = tcp_stream.set_nodelay(true);
270 #[cfg(target_os = "linux")]
271 {
272 let sock_ref = socket2::SockRef::from(&tcp_stream);
273 let _ = sock_ref.set_tcp_quickack(true);
274 }
275 let acceptor = acceptor.clone();
276 let state = state.clone();
277
278 let serve_fut = async move {
279 let tls_stream =
280 match tokio::time::timeout(TLS_HANDSHAKE_TIMEOUT, acceptor.accept(tcp_stream))
281 .await
282 {
283 Ok(Ok(stream)) => stream,
284 Ok(Err(e)) => {
285 tracing::debug!("[https] TLS handshake error: {}", e);
286 drop(permit);
287 return;
288 }
289 Err(_) => {
290 tracing::debug!("[https] TLS handshake timed out");
291 drop(permit);
292 return;
293 }
294 };
295
296 #[cfg(feature = "http2")]
299 let is_h2 = {
300 let (_, connection) = tls_stream.get_ref();
301 connection.alpn_protocol() == Some(b"h2")
302 };
303
304 let io = hyper_util::rt::TokioIo::new(tls_stream);
305 let svc = service_fn(move |req| hyper_handler(state.clone(), req, peer));
306
307 #[cfg(feature = "http2")]
308 if is_h2 {
309 let mut builder = hyper::server::conn::http2::Builder::new(LocalExecutor);
311 let _ = builder
312 .timer(hyper_util::rt::TokioTimer::new())
313 .initial_stream_window_size(65535)
314 .initial_connection_window_size(1024 * 1024)
315 .max_frame_size(16384)
316 .max_concurrent_streams(200)
317 .keep_alive_timeout(REQUEST_TIMEOUT);
318 #[cfg(feature = "ws")]
319 let _ = builder.enable_connect_protocol();
320
321 if let Err(e) = builder.serve_connection(io, svc).await {
322 tracing::debug!("[https] HTTP/2 Connection error: {}", e);
323 }
324 drop(permit);
325 return;
326 }
327
328 #[cfg(feature = "http1")]
335 {
336 let mut builder = hyper::server::conn::http1::Builder::new();
338 let _ = builder
339 .timer(hyper_util::rt::TokioTimer::new())
340 .header_read_timeout(REQUEST_TIMEOUT)
341 .keep_alive(true)
342 .max_buf_size(8192);
343
344 if let Err(e) = builder.serve_connection(io, svc).with_upgrades().await {
345 tracing::debug!("[https] HTTP/1.1 Connection error: {}", e);
346 }
347 }
348 drop(permit);
349 };
350
351 IS_LOCAL_WORKER.with(|flag| {
352 if flag.get() {
353 drop(tokio::task::spawn_local(serve_fut));
354 } else {
355 drop(tokio::spawn(serve_fut));
356 }
357 });
358 }
359 Ok(())
360 }
361
362 #[cfg(feature = "tls")]
370 pub async fn serve_https_config(
371 self,
372 listener: TcpListener,
373 config: rustls::ServerConfig,
374 ) -> Result<(), std::io::Error> {
375 crate::server::enforce_fips_compliance()?;
376 let acceptor = TlsAcceptor::from(Arc::new(config));
377 self.serve_https(listener, acceptor).await
378 }
379}
380
381pub(super) async fn hyper_handler<S>(
382 state: Arc<Server<S>>,
383 req: Request<hyper::body::Incoming>,
384 peer: std::net::SocketAddr,
385) -> Result<Response<Body>, std::io::Error>
386where
387 S: Clone + Send + Sync + 'static,
388{
389 let (parts, incoming_body) = req.into_parts();
390
391 let body = if HyperBody::is_end_stream(&incoming_body) {
392 Body::empty()
393 } else {
394 Body::stream(DeadlineBody {
395 inner: incoming_body,
396 deadline: None,
397 })
398 };
399
400 let mut rebuild_req = Request::from_parts(parts, body);
401 #[cfg(feature = "original-uri")]
402 {
403 let orig_uri = rebuild_req.uri().clone();
404 rebuild_req
405 .extensions_mut()
406 .insert(crate::routing::extract::OriginalUri(orig_uri));
407 }
408 rebuild_req
409 .extensions_mut()
410 .insert(crate::routing::extract::ConnectInfo(peer));
411 rebuild_req
412 .extensions_mut()
413 .insert(crate::routing::extract::MaxBodySize(state.max_body_size));
414
415 let resp = state.router.handle_request(rebuild_req).await;
416
417 if let Some((min, max)) = state.response_jitter {
418 tokio::time::sleep(crate::server::jittered_delay(min, max)).await;
419 }
420
421 Ok(resp)
422}