1use std::future::Future;
2use std::net::SocketAddr;
3#[cfg(feature = "tls")]
4use std::pin::Pin;
5use std::io::Write as _;
6use std::sync::{Arc, Mutex};
7use std::time::Duration;
8
9use hyper::header::{HeaderName, HeaderValue};
10use hyper::body::Body as HttpBody;
11use hyper::body::Bytes;
12use hyper::body::Incoming;
13use hyper::server::conn::http1;
14use hyper::{Method, Request, Response, StatusCode};
15use hyper::service::service_fn;
16use hyper_util::rt::TokioIo;
17use hyper_util::rt::TokioTimer;
18use http_body_util::combinators::BoxBody;
19use http_body_util::{BodyExt, Empty, Full};
20use tokio::net::{TcpListener, TcpStream};
21#[cfg(feature = "tls")]
22use tokio_rustls::TlsAcceptor;
23
24use crate::body::{MaxBodySize, DEFAULT_MAX_BODY_SIZE};
25use crate::cors::CorsConfig;
26use crate::error::ServeError;
27use crate::handler::{Handler, Middleware, OnUpgrade, ResponseBody};
28use crate::router::{QueryParams, Router};
29use crate::state::State;
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub struct PeerAddr(pub SocketAddr);
36
37const MAX_PATH_LEN: usize = 8_192;
38const MAX_QUERY_LEN: usize = 4_096;
39const DEFAULT_HEADER_READ_TIMEOUT: Duration = Duration::from_secs(30);
40const DEFAULT_MAX_CONNECTIONS: usize = 1024;
41
42const CONNECTION_OWNED_HEADERS: [HeaderName; 3] = [
47 hyper::header::CONTENT_LENGTH,
48 hyper::header::CONNECTION,
49 hyper::header::TRANSFER_ENCODING,
50];
51const DEFAULT_SHUTDOWN_DRAIN_TIMEOUT: Duration = Duration::from_secs(5);
59#[cfg(feature = "tls")]
60const DEFAULT_TLS_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
61const UPGRADE_HANDOFF_TIMEOUT: Duration = Duration::from_secs(10);
69const ACCEPT_BACKOFF_INITIAL: Duration = Duration::from_millis(10);
70const ACCEPT_BACKOFF_MAX: Duration = Duration::from_secs(1);
71
72#[cfg(test)]
73thread_local! {
74 static ERROR_LOG: std::cell::RefCell<Vec<(u16, String)>> = const { std::cell::RefCell::new(Vec::new()) };
75}
76
77#[cfg(test)]
78fn capture_error(code: u16, message: String) {
79 ERROR_LOG.with(|log| log.borrow_mut().push((code, message)));
80}
81
82#[cfg(test)]
83fn take_error_log() -> Vec<(u16, String)> {
84 ERROR_LOG.with(|log| log.borrow_mut().drain(..).collect())
85}
86
87trait TcpAccept {
92 async fn accept(&self) -> std::io::Result<(TcpStream, SocketAddr)>;
93}
94
95impl TcpAccept for TcpListener {
96 async fn accept(&self) -> std::io::Result<(TcpStream, SocketAddr)> {
97 TcpListener::accept(self).await
98 }
99}
100
101struct Backoff {
106 delay: Duration,
107}
108
109impl Backoff {
110 fn new() -> Self {
111 Backoff { delay: ACCEPT_BACKOFF_INITIAL }
112 }
113
114 fn next_delay(&mut self) -> Duration {
115 let delay = self.delay;
116 self.delay = (self.delay * 2).min(ACCEPT_BACKOFF_MAX);
117 delay
118 }
119
120 fn reset(&mut self) {
121 self.delay = ACCEPT_BACKOFF_INITIAL;
122 }
123}
124
125pub type ErrorHandler =
126 Arc<dyn Fn(StatusCode, &str) -> Response<ResponseBody> + Send + Sync>;
127
128pub(crate) type LogSink = Arc<Mutex<Box<dyn std::io::Write + Send>>>;
135
136pub struct App<S> {
165 state: Arc<S>,
166 log: Option<LogSink>,
167 extra_headers: Arc<Vec<(HeaderName, HeaderValue)>>,
168 router: Arc<Router<S>>,
169 max_body_size: usize,
170 pub(crate) upgrades_enabled: bool,
171 pub(crate) header_read_timeout: Duration,
172 pub(crate) max_connections: usize,
173 #[cfg(feature = "tls")]
174 pub(crate) tls_handshake_timeout: Duration,
175 error_handler: ErrorHandler,
176 cors_config: Option<CorsConfig>,
177}
178
179pub(crate) fn report_if_panicked(sink: &Option<LogSink>, joined: Result<(), tokio::task::JoinError>) {
185 let Err(e) = joined else {
186 return;
187 };
188 if e.is_panic() {
189 log_line(sink, format_args!("connection task panicked: {e}"));
190 }
191}
192
193pub(crate) fn log_line(sink: &Option<LogSink>, line: std::fmt::Arguments<'_>) {
199 let Some(sink) = sink else {
200 return;
201 };
202 if let Ok(mut out) = sink.lock() {
203 let _ = writeln!(out, "{line}");
204 let _ = out.flush();
205 }
206}
207
208fn parse_query(query: &str) -> QueryParams {
213 let mut map = std::collections::HashMap::new();
214 for pair in query.split('&').filter(|s| !s.is_empty()) {
215 if let Some((key, value)) = pair.split_once('=') {
216 let key = decode_query_component(key);
217 let value = decode_query_component(value);
218 map.insert(key, value);
219 } else {
220 let pair = decode_query_component(pair);
221 map.insert(pair, String::new());
222 }
223 }
224 QueryParams(map)
225}
226
227fn decode_query_component(s: &str) -> String {
228 let with_spaces = s.replace('+', " ");
229 percent_encoding::percent_decode_str(&with_spaces)
230 .decode_utf8_lossy()
231 .into_owned()
232}
233
234fn error_response(status: StatusCode, message: &str) -> Response<ResponseBody> {
235 #[cfg(test)]
236 capture_error(status.as_u16(), message.to_string());
237
238 let client_message = if status.is_server_error() {
239 "internal server error"
240 } else {
241 message
242 };
243
244 let body = serde_json::json!({ "message": client_message });
245 let json = serde_json::to_string(&body)
246 .unwrap_or_else(|_| r#"{"message":"internal server error"}"#.to_string());
247 let mut resp = Response::new(BoxBody::new(
248 Full::new(Bytes::from(json)).map_err(|never: std::convert::Infallible| match never {}),
249 ));
250 *resp.status_mut() = status;
251 resp.headers_mut().insert(
255 hyper::header::CONTENT_TYPE,
256 HeaderValue::from_static("application/json"),
257 );
258 resp
259}
260
261fn default_error_handler() -> ErrorHandler {
262 Arc::new(error_response)
263}
264
265fn ephemeral_bind_addr() -> SocketAddr {
269 (std::net::Ipv4Addr::LOCALHOST, 0).into()
270}
271
272impl<S: Send + Sync + 'static> App<S> {
273 pub fn new(state: S) -> Self {
278 App {
279 state: Arc::new(state),
280 router: Arc::new(Router::new()),
281 max_body_size: DEFAULT_MAX_BODY_SIZE,
282 header_read_timeout: DEFAULT_HEADER_READ_TIMEOUT,
283 upgrades_enabled: false,
284 log: None,
285 extra_headers: Arc::new(Vec::new()),
286 max_connections: DEFAULT_MAX_CONNECTIONS,
287 #[cfg(feature = "tls")]
288 tls_handshake_timeout: DEFAULT_TLS_HANDSHAKE_TIMEOUT,
289 error_handler: default_error_handler(),
290 cors_config: None,
291 }
292 }
293
294 pub fn state_arc(&self) -> Arc<S> {
299 Arc::clone(&self.state)
300 }
301
302 pub async fn route(&self, req: Request<Incoming>) -> Response<ResponseBody> {
303 self.route_with(req, None).await
304 }
305
306 pub async fn route_with_peer(&self, req: Request<Incoming>, peer: SocketAddr) -> Response<ResponseBody> {
313 self.route_with(req, Some(peer)).await
314 }
315
316 async fn route_with(&self, req: Request<Incoming>, peer: Option<SocketAddr>) -> Response<ResponseBody> {
317 let method = req.method().clone();
318 let req_origin = req
323 .headers()
324 .get(hyper::header::ORIGIN)
325 .and_then(|v| v.to_str().ok())
326 .map(|s| s.to_string());
327
328 let mut resp = self.route_inner(req, peer, req_origin.as_deref()).await;
329 self.finalize(&mut resp, req_origin.as_deref());
330
331 if method == Method::HEAD {
338 let (mut parts, body) = resp.into_parts();
339 if !parts.headers.contains_key(hyper::header::CONTENT_LENGTH) {
347 if let Some(len) = HttpBody::size_hint(&body).exact() {
348 parts.headers.insert(hyper::header::CONTENT_LENGTH, len.into());
349 }
350 }
351 Response::from_parts(parts, BoxBody::new(Empty::new().map_err(|never: std::convert::Infallible| match never {})))
352 } else {
353 resp
354 }
355 }
356
357 fn finalize(&self, resp: &mut Response<ResponseBody>, req_origin: Option<&str>) {
369 if let Some(cfg) = &self.cors_config {
370 cfg.apply_to_response(resp, req_origin);
371 }
372
373 let headers = resp.headers_mut();
376 headers
377 .entry(hyper::header::X_CONTENT_TYPE_OPTIONS)
378 .or_insert(HeaderValue::from_static("nosniff"));
379
380 for (name, value) in self.extra_headers.iter() {
381 headers.entry(name).or_insert(value.clone());
382 }
383 }
384
385 async fn route_inner(
390 &self,
391 req: Request<Incoming>,
392 peer: Option<SocketAddr>,
393 req_origin: Option<&str>,
394 ) -> Response<ResponseBody> {
395 if req.version() == hyper::Version::HTTP_11
402 && req.uri().authority().is_none()
403 && !req.headers().contains_key(hyper::header::HOST)
404 {
405 return (self.error_handler)(StatusCode::BAD_REQUEST, "missing host header");
406 }
407 if req.uri().path().len() > MAX_PATH_LEN {
408 return (self.error_handler)(StatusCode::BAD_REQUEST, "path too long");
409 }
410 if req.uri().query().map(|q| q.len()).unwrap_or(0) > MAX_QUERY_LEN {
411 return (self.error_handler)(StatusCode::BAD_REQUEST, "query string too long");
412 }
413
414 let method = req.method().clone();
415 let path = req.uri().path().to_string();
416 let state = State::from_arc(Arc::clone(&self.state));
417 let query = req.uri().query().unwrap_or("");
424 let query_params = if query.is_empty() {
425 QueryParams::default()
426 } else {
427 parse_query(query)
428 };
429
430 if method == Method::OPTIONS && req_origin.is_some() {
432 if let Some(cfg) = &self.cors_config {
433 if self.router.path_exists(&path) {
434 let requested_headers = req
435 .headers()
436 .get("access-control-request-headers")
437 .and_then(|v| v.to_str().ok());
438 let allowed = self.allowed_methods_with_head(&path);
439 return cfg.preflight_response(req_origin, requested_headers, &allowed);
440 }
441 }
442 }
443
444 let method_to_match = if method == Method::HEAD {
445 Method::GET
446 } else {
447 method.clone()
448 };
449
450 match self.router.match_route(&method_to_match, &path) {
451 Some((handler, params)) => {
452 let mut req = req;
453 if !query_params.0.is_empty() {
459 req.extensions_mut().insert(query_params);
460 }
461 if !params.0.is_empty() {
462 req.extensions_mut().insert(params);
463 }
464 if self.max_body_size != DEFAULT_MAX_BODY_SIZE {
468 req.extensions_mut().insert(MaxBodySize(self.max_body_size));
469 }
470 if let Some(peer) = peer {
471 req.extensions_mut().insert(PeerAddr(peer));
472 }
473 match handler(req, state).await {
474 Ok(resp) => resp,
475 Err(e) => {
476 let status = StatusCode::from_u16(e.code)
477 .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR);
478 if status.is_server_error() {
484 log_line(&self.log, format_args!("{status} {}: {}", path, e.message));
485 }
486 (self.error_handler)(status, &e.message)
487 }
488 }
489 }
490 None => {
491 let allowed = self.allowed_methods_with_head(&path);
492 if !allowed.is_empty() {
493 let mut method_strs: Vec<&str> = allowed.iter().map(|m| m.as_str()).collect();
494 method_strs.sort();
495 method_strs.dedup();
496 let allow_header = method_strs.join(", ");
497 let mut resp = (self.error_handler)(StatusCode::METHOD_NOT_ALLOWED, "method not allowed");
498 if let Ok(val) = allow_header.parse() {
499 resp.headers_mut().insert("allow", val);
500 }
501 resp
502 } else {
503 (self.error_handler)(StatusCode::NOT_FOUND, "not found")
504 }
505 }
506 }
507 }
508
509 fn allowed_methods_with_head(&self, path: &str) -> Vec<Method> {
516 let mut allowed = self.router.allowed_methods(path);
517 if allowed.contains(&Method::GET) {
518 allowed.push(Method::HEAD);
519 }
520 allowed
521 }
522
523 pub async fn bind_ephemeral(self) -> Result<u16, ServeError> {
529 let listener = TcpListener::bind(ephemeral_bind_addr())
530 .await
531 .map_err(|e| ServeError::new(500, format!("failed to bind to ephemeral port: {e}")))?;
532 let port = listener
533 .local_addr()
534 .map_err(|e| ServeError::new(500, format!("failed to get assigned port: {e}")))?
535 .port();
536 let app = Arc::new(self);
537 tokio::spawn(async move {
538 serve_loop(listener, app, std::future::pending(), plain_connect).await;
539 });
540 Ok(port)
541 }
542
543 #[cfg(feature = "tls")]
549 pub async fn bind_tls_ephemeral(
550 self,
551 config: Arc<rustls::ServerConfig>,
552 ) -> Result<u16, ServeError> {
553 let listener = TcpListener::bind(ephemeral_bind_addr())
554 .await
555 .map_err(|e| ServeError::new(500, format!("failed to bind to ephemeral port: {e}")))?;
556 let port = listener
557 .local_addr()
558 .map_err(|e| ServeError::new(500, format!("failed to get assigned port: {e}")))?
559 .port();
560 let handshake_timeout = self.tls_handshake_timeout;
561 let acceptor = TlsAcceptor::from(config);
562 let app = Arc::new(self);
563 tokio::spawn(async move {
564 serve_loop(listener, app, std::future::pending(), tls_connect(acceptor, handshake_timeout)).await;
565 });
566 Ok(port)
567 }
568
569 pub async fn run<F>(self, listener: TcpListener, shutdown: F) -> Result<(), ServeError>
574 where
575 F: Future<Output = ()> + Send + 'static,
576 {
577 let app = Arc::new(self);
578 serve_loop(listener, app, shutdown, plain_connect).await;
579 Ok(())
580 }
581
582 pub async fn bind(self, addr: SocketAddr) -> Result<(), ServeError> {
587 let shutdown = signal_shutdown()?;
588 let listener = TcpListener::bind(addr)
589 .await
590 .map_err(|e| ServeError::new(500, format!("failed to bind to {addr}: {e}")))?;
591 self.run(listener, shutdown).await
592 }
593
594 #[cfg(feature = "tls")]
596 pub async fn run_tls<F>(
597 self,
598 listener: TcpListener,
599 config: Arc<rustls::ServerConfig>,
600 shutdown: F,
601 ) -> Result<(), ServeError>
602 where
603 F: Future<Output = ()> + Send + 'static,
604 {
605 let handshake_timeout = self.tls_handshake_timeout;
606 let acceptor = TlsAcceptor::from(config);
607 let app = Arc::new(self);
608 serve_loop(listener, app, shutdown, tls_connect(acceptor, handshake_timeout)).await;
609 Ok(())
610 }
611
612 #[cfg(feature = "tls")]
614 pub async fn bind_tls(
615 self,
616 addr: SocketAddr,
617 config: Arc<rustls::ServerConfig>,
618 ) -> Result<(), ServeError> {
619 let shutdown = signal_shutdown()?;
620 let listener = TcpListener::bind(addr)
621 .await
622 .map_err(|e| ServeError::new(500, format!("failed to bind to {addr}: {e}")))?;
623 self.run_tls(listener, config, shutdown).await
624 }
625}
626
627impl App<()> {
628 pub fn stateless() -> Self {
629 App::new(())
630 }
631}
632
633async fn serve_connection<S, IO>(io: IO, app: Arc<App<S>>, header_read_timeout: Duration, peer: SocketAddr)
637where
638 S: Send + Sync + 'static,
639 IO: hyper::rt::Read + hyper::rt::Write + Unpin + Send + 'static,
640{
641 let app_for_conn = app.clone();
642 let log_for_conn = app.log.clone();
643
644 let pending_upgrade: Arc<Mutex<Option<(hyper::upgrade::OnUpgrade, OnUpgrade)>>> =
649 Arc::new(Mutex::new(None));
650 let pending_for_service = pending_upgrade.clone();
651
652 let svc = service_fn(move |mut req: Request<Incoming>| {
653 let app = app.clone();
654 let pending = pending_for_service.clone();
655 let upgrade = hyper::upgrade::on(&mut req);
658 async move {
659 let observed = app.log.as_ref().map(|_| {
661 (
662 std::time::Instant::now(),
663 req.method().clone(),
664 req.uri().path().to_string(),
665 )
666 });
667
668 let resp = app.route_with_peer(req, peer).await;
669
670 if let Some((started, method, path)) = observed {
671 log_line(
672 &app.log,
673 format_args!(
674 "{method} {path} {} {:.3}ms",
675 resp.status().as_u16(),
676 started.elapsed().as_secs_f64() * 1000.0
677 ),
678 );
679 }
680 let mut resp = resp;
681 if let Some(callback) = resp.extensions_mut().remove::<OnUpgrade>() {
682 if app.upgrades_enabled {
683 *pending.lock().unwrap() = Some((upgrade, callback));
684 } else {
685 log_line(
688 &app.log,
689 format_args!(
690 "handler returned an upgrade for {} but the app was not built \
691 with with_upgrades(); the connection will not be upgraded",
692 resp.status().as_u16()
693 ),
694 );
695 }
696 }
697 Ok::<_, hyper::Error>(resp)
698 }
699 });
700 let mut builder = http1::Builder::new();
701 builder.timer(TokioTimer::new());
702 builder.header_read_timeout(header_read_timeout);
703
704 if app_for_conn.upgrades_enabled {
705 let _ = builder.serve_connection(io, svc).with_upgrades().await;
706 } else {
707 let _ = builder.serve_connection(io, svc).await;
708 }
709
710 let taken = pending_upgrade.lock().unwrap().take();
714 if let Some((upgrade, callback)) = taken {
715 match tokio::time::timeout(UPGRADE_HANDOFF_TIMEOUT, upgrade).await {
716 Ok(Ok(upgraded)) => callback.run(TokioIo::new(upgraded)).await,
717 Ok(Err(e)) => log_line(&log_for_conn, format_args!("upgrade failed: {e}")),
718 Err(_) => log_line(
719 &log_for_conn,
720 format_args!(
721 "upgrade did not complete within {UPGRADE_HANDOFF_TIMEOUT:?}; \
722 releasing the connection"
723 ),
724 ),
725 }
726 }
727}
728
729async fn plain_connect(stream: TcpStream) -> Option<TcpStream> {
730 Some(stream)
731}
732
733#[cfg(feature = "tls")]
738fn tls_connect(
739 acceptor: TlsAcceptor,
740 handshake_timeout: Duration,
741) -> impl Fn(TcpStream) -> Pin<Box<dyn Future<Output = Option<tokio_rustls::server::TlsStream<TcpStream>>> + Send>> + Clone {
742 move |stream| {
743 let acceptor = acceptor.clone();
744 Box::pin(async move {
745 match tokio::time::timeout(handshake_timeout, acceptor.accept(stream)).await {
746 Ok(Ok(s)) => Some(s),
747 Ok(Err(_)) | Err(_) => None,
748 }
749 })
750 }
751}
752
753async fn accept_and_permit<L: TcpAccept>(
767 listener: &L,
768 backoff: &mut Backoff,
769 semaphore: &Arc<tokio::sync::Semaphore>,
770) -> Option<(TcpStream, SocketAddr, tokio::sync::OwnedSemaphorePermit)> {
771 let permit = semaphore.clone().acquire_owned().await.ok()?;
781
782 loop {
783 match listener.accept().await {
784 Ok((stream, peer)) => {
785 backoff.reset();
786 return Some((stream, peer, permit));
787 }
788 Err(_) => tokio::time::sleep(backoff.next_delay()).await,
789 }
790 }
791}
792
793fn signal_install_error(signal: &str, cause: std::io::Error) -> ServeError {
799 ServeError::new(500, format!("failed to install {signal} handler: {cause}"))
800}
801
802fn signal_shutdown() -> Result<impl Future<Output = ()> + Send, ServeError> {
808 let mut sigint = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::interrupt())
809 .map_err(|e| signal_install_error("SIGINT", e))?;
810 let mut sigterm = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
811 .map_err(|e| signal_install_error("SIGTERM", e))?;
812
813 Ok(async move {
814 tokio::select! {
815 _ = sigint.recv() => {}
816 _ = sigterm.recv() => {}
817 }
818 })
819}
820
821async fn serve_loop<S, F, C, Fut, IO>(listener: TcpListener, app: Arc<App<S>>, shutdown: F, connect: C)
834where
835 S: Send + Sync + 'static,
836 F: Future<Output = ()> + Send + 'static,
837 C: Fn(TcpStream) -> Fut + Send + Sync + 'static,
838 Fut: Future<Output = Option<IO>> + Send + 'static,
839 IO: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static,
840{
841 let header_read_timeout = app.header_read_timeout;
842 let semaphore = Arc::new(tokio::sync::Semaphore::new(app.max_connections));
843 let connect = Arc::new(connect);
844 let mut backoff = Backoff::new();
845 let mut join_set: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
846 let mut shutdown_pin = std::pin::pin!(shutdown);
847 let mut shutting_down = false;
848
849 loop {
850 if !shutting_down {
851 tokio::select! {
852 accepted = accept_and_permit(&listener, &mut backoff, &semaphore) => {
853 match accepted {
854 Some((stream, peer, permit)) => {
855 let app = app.clone();
856 let connect = connect.clone();
857 join_set.spawn(async move {
858 let _permit = permit;
859 if let Some(io) = connect(stream).await {
860 serve_connection(TokioIo::new(io), app, header_read_timeout, peer).await;
861 }
862 });
863 }
864 None => shutting_down = true,
865 }
866 }
867 joined = join_set.join_next(), if !join_set.is_empty() => {
872 if let Some(joined) = joined {
873 report_if_panicked(&app.log, joined);
874 }
875 }
876 _ = shutdown_pin.as_mut() => {
877 shutting_down = true;
878 }
879 }
880 continue;
881 }
882
883 let drained = tokio::time::timeout(DEFAULT_SHUTDOWN_DRAIN_TIMEOUT, async {
890 while let Some(joined) = join_set.join_next().await {
891 report_if_panicked(&app.log, joined);
892 }
893 })
894 .await;
895
896 if drained.is_err() {
897 join_set.shutdown().await;
898 }
899 break;
900 }
901}
902
903#[must_use = "RouteBuilder does nothing until .seal() is called"]
928pub struct RouteBuilder<S> {
929 state: Arc<S>,
930 router: Router<S>,
931 max_body_size: usize,
932 header_read_timeout: Duration,
933 log: Option<LogSink>,
934 extra_headers: Arc<Vec<(HeaderName, HeaderValue)>>,
935 max_connections: usize,
936 #[cfg(feature = "tls")]
937 tls_handshake_timeout: Duration,
938 error_handler: ErrorHandler,
939 cors_config: Option<CorsConfig>,
940 middlewares: Vec<Middleware<S>>,
941 upgrades_enabled: bool,
942}
943
944impl<S: Send + Sync + 'static> RouteBuilder<S> {
945 pub fn new(state: S) -> Self {
947 RouteBuilder {
948 state: Arc::new(state),
949 router: Router::new(),
950 max_body_size: DEFAULT_MAX_BODY_SIZE,
951 header_read_timeout: DEFAULT_HEADER_READ_TIMEOUT,
952 upgrades_enabled: false,
953 log: None,
954 extra_headers: Arc::new(Vec::new()),
955 max_connections: DEFAULT_MAX_CONNECTIONS,
956 #[cfg(feature = "tls")]
957 tls_handshake_timeout: DEFAULT_TLS_HANDSHAKE_TIMEOUT,
958 error_handler: default_error_handler(),
959 cors_config: None,
960 middlewares: Vec::new(),
961 }
962 }
963
964 pub fn wrap(mut self, middleware: Middleware<S>) -> Self {
968 self.middlewares.push(middleware);
969 self
970 }
971
972 fn apply_middlewares(&self, handler: Handler<S>) -> Handler<S> {
973 self.middlewares.iter().rev().fold(handler, |acc, mw| mw(acc))
974 }
975
976 pub fn with_max_body_size(mut self, max: usize) -> Self {
978 self.max_body_size = max;
979 self
980 }
981
982 pub fn with_header_read_timeout(mut self, d: Duration) -> Self {
984 self.header_read_timeout = d;
985 self
986 }
987
988 pub fn with_response_header(mut self, name: &str, value: &str) -> Result<Self, ServeError> {
1011 let name = HeaderName::from_bytes(name.as_bytes())
1012 .map_err(|_| ServeError::new(500, format!("invalid header name: {name}")))?;
1013 let value = HeaderValue::from_str(value)
1014 .map_err(|_| ServeError::new(500, format!("invalid value for header {name}")))?;
1015
1016 if CONNECTION_OWNED_HEADERS.contains(&name) {
1017 return Err(ServeError::new(
1018 500,
1019 format!("{name} is owned by the connection layer and cannot be set as a fixed header"),
1020 ));
1021 }
1022
1023 Arc::make_mut(&mut self.extra_headers).push((name, value));
1024 Ok(self)
1025 }
1026
1027 pub fn with_request_logging(self) -> Self {
1028 self.with_request_logging_to(Box::new(std::io::stderr()))
1029 }
1030
1031 pub fn with_request_logging_to(mut self, writer: Box<dyn std::io::Write + Send>) -> Self {
1048 self.log = Some(Arc::new(Mutex::new(writer)));
1049 self
1050 }
1051
1052 #[cfg(feature = "tls")]
1055 pub fn with_tls_handshake_timeout(mut self, d: Duration) -> Self {
1056 self.tls_handshake_timeout = d;
1057 self
1058 }
1059
1060 pub fn with_upgrades(mut self) -> Self {
1074 self.upgrades_enabled = true;
1075 self
1076 }
1077
1078 pub fn with_max_connections(mut self, max: usize) -> Self {
1079 self.max_connections = max;
1080 self
1081 }
1082
1083 pub fn with_error_handler(
1084 mut self,
1085 f: impl Fn(StatusCode, &str) -> Response<ResponseBody> + Send + Sync + 'static,
1086 ) -> Self {
1087 self.error_handler = Arc::new(f);
1088 self
1089 }
1090
1091 pub fn with_cors(mut self, config: CorsConfig) -> Self {
1092 self.cors_config = Some(config);
1093 self
1094 }
1095
1096 pub fn get(mut self, path: &str, handler: Handler<S>) -> Self {
1097 let handler = self.apply_middlewares(handler);
1098 self.router.insert(Method::GET, path, handler);
1099 self
1100 }
1101
1102 pub fn post(mut self, path: &str, handler: Handler<S>) -> Self {
1103 let handler = self.apply_middlewares(handler);
1104 self.router.insert(Method::POST, path, handler);
1105 self
1106 }
1107
1108 pub fn put(mut self, path: &str, handler: Handler<S>) -> Self {
1109 let handler = self.apply_middlewares(handler);
1110 self.router.insert(Method::PUT, path, handler);
1111 self
1112 }
1113
1114 pub fn delete(mut self, path: &str, handler: Handler<S>) -> Self {
1115 let handler = self.apply_middlewares(handler);
1116 self.router.insert(Method::DELETE, path, handler);
1117 self
1118 }
1119
1120 pub fn patch(mut self, path: &str, handler: Handler<S>) -> Self {
1125 let handler = self.apply_middlewares(handler);
1126 self.router.insert(Method::PATCH, path, handler);
1127 self
1128 }
1129
1130 pub fn seal(self) -> App<S> {
1131 App {
1132 state: self.state,
1133 router: Arc::new(self.router),
1134 max_body_size: self.max_body_size,
1135 header_read_timeout: self.header_read_timeout,
1136 upgrades_enabled: self.upgrades_enabled,
1137 log: self.log,
1138 extra_headers: self.extra_headers,
1139 max_connections: self.max_connections,
1140 #[cfg(feature = "tls")]
1141 tls_handshake_timeout: self.tls_handshake_timeout,
1142 error_handler: self.error_handler,
1143 cors_config: self.cors_config,
1144 }
1145 }
1146}
1147
1148impl RouteBuilder<()> {
1149 pub fn stateless() -> Self {
1150 RouteBuilder::new(())
1151 }
1152}
1153
1154#[cfg(test)]
1155#[path = "../tests/unit/app.rs"]
1156mod tests;