1use std::future::{Future, IntoFuture};
4use std::net::{IpAddr, SocketAddr};
5use std::str::FromStr;
6
7use axum::extract::DefaultBodyLimit;
8use axum::http::StatusCode;
9use axum::response::{IntoResponse, Response};
10use axum::routing::get;
11use axum::{Router, serve};
12use tokio::net::TcpListener;
13use tower_http::LatencyUnit;
14use tower_http::trace::{DefaultMakeSpan, DefaultOnResponse, TraceLayer};
15use tracing::{Level, error, info, instrument, warn};
16
17use crate::env;
18use crate::environment::Environment;
19
20use super::error::WebServerError;
21use super::shutdown::{AppShutdown, DEFAULT_DRAIN_TIMEOUT, Shutdown};
22use super::state::{StateParts, WebServerState};
23
24pub const ENV_HOST: &str = "WSB_HOST";
26
27pub const ENV_PORT: &str = "WSB_PORT";
29
30pub const DEFAULT_PORT: u16 = 8080;
32
33pub const DEFAULT_BODY_LIMIT: usize = 256 * 1024;
38
39pub const API_PREFIX: &str = "/api/v1";
45
46pub struct WebServer<S = ()> {
53 host: String,
54 port: u16,
55 environment: Environment,
56 app: S,
57
58 body_limit: usize,
59 router: Router<WebServerState<S>>,
60
61 #[cfg(feature = "pages")]
62 frontend: Option<super::frontend::FrontendParams<S>>,
63}
64
65impl WebServer<()> {
66 #[must_use]
68 pub fn new(host: impl Into<String>, port: u16, environment: Environment) -> Self {
69 Self::with_state(host, port, environment, ())
70 }
71
72 pub fn from_env() -> Result<Self, WebServerError> {
82 Self::from_env_with_state(())
83 }
84}
85
86impl<S> WebServer<S>
87where
88 S: Clone + Send + Sync + 'static,
89{
90 #[must_use]
92 pub fn with_state(
93 host: impl Into<String>,
94 port: u16,
95 environment: Environment,
96 app: S,
97 ) -> Self {
98 Self {
99 host: host.into(),
100 port,
101 environment,
102 app,
103 body_limit: DEFAULT_BODY_LIMIT,
104 router: Router::new(),
105 #[cfg(feature = "pages")]
106 frontend: None,
107 }
108 }
109
110 pub fn from_env_with_state(app: S) -> Result<Self, WebServerError> {
117 let environment: Environment = Environment::from_env()?;
118 let host: String = env::optional(ENV_HOST).unwrap_or_else(|| {
119 if environment.is_production() {
120 String::from("0.0.0.0")
121 } else {
122 String::from("127.0.0.1")
123 }
124 });
125 let port: u16 =
126 env::parse_or(ENV_PORT, "port number", DEFAULT_PORT).map_err(WebServerError::Env)?;
127
128 Ok(Self::with_state(host, port, environment, app))
129 }
130
131 #[must_use]
136 pub const fn body_limit(mut self, body_limit: usize) -> Self {
137 self.body_limit = body_limit;
138 self
139 }
140
141 #[must_use]
143 pub fn nest(mut self, path: &str, router: Router<WebServerState<S>>) -> Self {
144 self.router = self.router.nest(path, router);
145 self
146 }
147
148 #[must_use]
150 pub fn merge(mut self, router: Router<WebServerState<S>>) -> Self {
151 self.router = self.router.merge(router);
152 self
153 }
154
155 #[must_use]
157 pub fn nest_service<T>(mut self, path: &str, service: T) -> Self
158 where
159 T: tower::Service<axum::extract::Request, Error = std::convert::Infallible>
160 + Clone
161 + Send
162 + Sync
163 + 'static,
164 T::Response: axum::response::IntoResponse,
165 T::Future: Send + 'static,
166 {
167 self.router = self.router.nest_service(path, service);
168 self
169 }
170
171 #[cfg(feature = "pages")]
178 #[must_use]
179 pub fn frontend(mut self, params: super::frontend::FrontendParams<S>) -> Self {
180 self.frontend = Some(params);
181 self
182 }
183
184 #[instrument(skip_all)]
191 pub async fn run(self, shutdown: Shutdown) -> Result<(), WebServerError>
192 where
193 S: AppShutdown,
194 {
195 let Self {
196 host,
197 port,
198 environment,
199 app,
200 body_limit,
201 router,
202 #[cfg(feature = "pages")]
203 frontend,
204 } = self;
205
206 let cache_buster: crate::assets::CacheBuster = crate::assets::CacheBuster::load()?;
209
210 let mut no_cache: Router<WebServerState<S>> = router;
211 let mut built_in: Router<WebServerState<S>> = Router::new().route("/health", get(health));
212
213 #[cfg(feature = "pages")]
214 let mut frontend_runtime: Option<crate::templates::FrontendRuntime> = None;
215 #[cfg(feature = "pages")]
216 let mut base: Option<crate::templates::BaseTemplateData> = None;
217 #[cfg(feature = "pages")]
218 let mut templates: Option<crate::templates::TemplateRegistry<'static>> = None;
219 #[cfg(feature = "pages")]
220 let mut not_found: Option<(
221 std::sync::Arc<crate::templates::PageTemplateData>,
222 std::sync::Arc<serde_json::Value>,
223 )> = None;
224 #[cfg(feature = "pages")]
225 let mut proxy_scripts: Option<Router<WebServerState<S>>> = None;
226
227 #[cfg(feature = "pages")]
228 if let Some(params) = frontend {
229 let registry: crate::templates::TemplateRegistry<'static> =
230 crate::templates::TemplateRegistry::from_dir(crate::templates::TEMPLATE_ROOT)?;
231
232 let built: super::frontend::Frontend<S> =
233 super::frontend::Frontend::build(params, &cache_buster, environment)?;
234
235 no_cache = no_cache.merge(well_known_routes(&built.well_known));
236 no_cache = no_cache.merge(icon_routes(&cache_buster, built.has_svg_icon));
237
238 let (scripts, endpoints) = proxy_routes(&built);
239 proxy_scripts = Some(scripts);
243 built_in = built_in.merge(endpoints);
244
245 not_found = Some(built.not_found.clone());
246 no_cache = no_cache.merge(built.pages.into_router());
247
248 frontend_runtime = Some(built.runtime);
249 base = Some(built.base);
250 templates = Some(registry);
251 }
252
253 let no_cache: Router<WebServerState<S>> = no_cache.nest(API_PREFIX, built_in);
254
255 let mut app_router: Router<WebServerState<S>> = apply_cache_policy(no_cache, &cache_buster);
256
257 #[cfg(feature = "pages")]
261 if let Some(scripts) = proxy_scripts {
262 app_router = app_router.merge(scripts);
263 }
264
265 #[cfg(feature = "pages")]
266 let app_router: Router<WebServerState<S>> = attach_not_found(app_router, not_found);
267 #[cfg(not(feature = "pages"))]
268 let app_router: Router<WebServerState<S>> = app_router.fallback(plain_not_found);
269
270 let state: WebServerState<S> = WebServerState::new(StateParts {
271 host: host.clone(),
272 port,
273 environment,
274 shutdown: shutdown.clone(),
275 #[cfg(feature = "templates")]
276 base,
277 #[cfg(feature = "templates")]
278 templates,
279 cache_buster: Some(cache_buster),
280 #[cfg(feature = "templates")]
281 frontend: frontend_runtime,
282 app,
283 });
284
285 let cleanup: WebServerState<S> = state.clone();
288
289 let app_router: Router = app_router
292 .with_state(state)
293 .layer(
294 TraceLayer::new_for_http()
295 .make_span_with(
301 DefaultMakeSpan::new()
302 .level(Level::INFO)
303 .include_headers(false),
304 )
305 .on_response(
306 DefaultOnResponse::new()
307 .level(Level::INFO)
308 .latency_unit(LatencyUnit::Millis),
309 ),
310 )
311 .layer(DefaultBodyLimit::max(body_limit));
312
313 serve_on(app_router, &host, port, shutdown, async move {
314 cleanup.app().on_shutdown().await;
315 })
316 .await
317 }
318}
319
320#[cfg(feature = "pages")]
327fn attach_not_found<S>(
328 router: Router<WebServerState<S>>,
329 not_found: Option<(
330 std::sync::Arc<crate::templates::PageTemplateData>,
331 std::sync::Arc<serde_json::Value>,
332 )>,
333) -> Router<WebServerState<S>>
334where
335 S: Clone + Send + Sync + 'static,
336{
337 match not_found {
338 Some((page, data)) => router.fallback(move |axum::extract::State(state)| {
339 let page: std::sync::Arc<crate::templates::PageTemplateData> =
340 std::sync::Arc::clone(&page);
341 let data: std::sync::Arc<serde_json::Value> = std::sync::Arc::clone(&data);
342 async move {
343 let body: Response = super::pages::render_or_500(&state, &page, &data);
344 (StatusCode::NOT_FOUND, body).into_response()
345 }
346 }),
347 None => router.fallback(plain_not_found),
348 }
349}
350
351async fn serve_on<F>(
358 router: Router,
359 host: &str,
360 port: u16,
361 shutdown: Shutdown,
362 on_shutdown: F,
363) -> Result<(), WebServerError>
364where
365 F: Future<Output = ()> + Send,
366{
367 let ip: IpAddr = IpAddr::from_str(host).map_err(|source| WebServerError::Bind {
368 addr: SocketAddr::from(([0, 0, 0, 0], port)),
369 source: std::io::Error::new(std::io::ErrorKind::InvalidInput, source),
370 })?;
371 let address: SocketAddr = SocketAddr::new(ip, port);
372 let listener: TcpListener =
373 TcpListener::bind(address)
374 .await
375 .map_err(|source| WebServerError::Bind {
376 addr: address,
377 source,
378 })?;
379
380 info!("listening on http://{address}");
381
382 let serving = serve(
386 listener,
387 router.into_make_service_with_connect_info::<SocketAddr>(),
388 )
389 .with_graceful_shutdown(shutdown.clone().recv())
390 .into_future();
391 tokio::pin!(serving);
392
393 tokio::select! {
397 result = &mut serving => return result.map_err(WebServerError::Serve),
398 () = shutdown.recv() => {}
399 }
400
401 let (served, cleaned) = tokio::join!(
406 tokio::time::timeout(DEFAULT_DRAIN_TIMEOUT, &mut serving),
407 tokio::time::timeout(DEFAULT_DRAIN_TIMEOUT, on_shutdown),
408 );
409
410 if cleaned.is_err() {
414 error!(
415 "app shutdown hook exceeded {DEFAULT_DRAIN_TIMEOUT:?}; cleanup was cancelled part-way"
416 );
417 }
418
419 match served {
420 Ok(result) => result.map_err(WebServerError::Serve),
421 Err(_elapsed) => {
424 warn!("drain window elapsed after {DEFAULT_DRAIN_TIMEOUT:?}; closing what remained");
425 Ok(())
426 }
427 }
428}
429
430fn apply_cache_policy<S>(
433 router: Router<WebServerState<S>>,
434 cache_buster: &crate::assets::CacheBuster,
435) -> Router<WebServerState<S>>
436where
437 S: Clone + Send + Sync + 'static,
438{
439 let router: Router<WebServerState<S>> = router.layer(axum::middleware::from_fn(
440 crate::assets::CacheBuster::never_cache_middleware,
441 ));
442
443 if cache_buster.is_empty() {
444 return router;
445 }
446
447 router.merge(
448 Router::new()
449 .nest_service(
450 "/static",
451 tower_http::services::ServeDir::new(crate::assets::STATIC_DIRECTORY),
452 )
453 .layer(axum::middleware::from_fn(
454 crate::assets::CacheBuster::forever_cache_middleware,
455 )),
456 )
457}
458
459async fn health() -> StatusCode {
462 StatusCode::OK
463}
464
465async fn plain_not_found() -> Response {
467 (StatusCode::NOT_FOUND, "404").into_response()
468}
469
470#[cfg(feature = "pages")]
477fn well_known_routes<S>(well_known: &super::frontend::WellKnown) -> Router<WebServerState<S>>
478where
479 S: Clone + Send + Sync + 'static,
480{
481 use axum::http::header;
482
483 fn text<S>(
484 router: Router<WebServerState<S>>,
485 path: &str,
486 content_type: &'static str,
487 body: String,
488 ) -> Router<WebServerState<S>>
489 where
490 S: Clone + Send + Sync + 'static,
491 {
492 router.route(
493 path,
494 get(move || {
495 let body: String = body.clone();
496 async move { ([(header::CONTENT_TYPE, content_type)], body) }
497 }),
498 )
499 }
500
501 let mut router: Router<WebServerState<S>> = Router::new();
502 router = text(
503 router,
504 "/robots.txt",
505 "text/plain; charset=utf-8",
506 well_known.robots_txt.clone(),
507 );
508 router = text(
509 router,
510 "/humans.txt",
511 "text/plain; charset=utf-8",
512 well_known.humans_txt.clone(),
513 );
514 router = text(
515 router,
516 "/site.webmanifest",
517 "application/manifest+json",
518 well_known.webmanifest.clone(),
519 );
520 router = text(
521 router,
522 crate::sitemap::SITEMAP_INDEX_PATH,
523 "application/xml",
524 well_known.sitemaps.index().to_string(),
525 );
526 for (index, chunk) in well_known.sitemaps.chunks().iter().enumerate() {
527 router = text(
528 router,
529 &format!("/sitemap-{}.xml", index + 1),
530 "application/xml",
531 chunk.clone(),
532 );
533 }
534 router
535}
536
537#[cfg(feature = "pages")]
543fn icon_routes<S>(
544 cache_buster: &crate::assets::CacheBuster,
545 has_svg_icon: bool,
546) -> Router<WebServerState<S>>
547where
548 S: Clone + Send + Sync + 'static,
549{
550 use tower_http::services::ServeFile;
551
552 let mut icons: Vec<(&str, String)> = vec![
553 (
554 "/favicon.ico",
555 String::from("static/image/favicon/favicon.ico"),
556 ),
557 (
558 "/apple-touch-icon.png",
559 String::from("static/image/favicon/apple-touch-icon.png"),
560 ),
561 (
562 "/icon-192.png",
563 String::from("static/image/favicon/icon-192.png"),
564 ),
565 (
566 "/icon-512.png",
567 String::from("static/image/favicon/icon-512.png"),
568 ),
569 ];
570 if has_svg_icon {
572 icons.push((
573 "/favicon.svg",
574 String::from("static/image/favicon/favicon.svg"),
575 ));
576 }
577
578 let mut router: Router<WebServerState<S>> = Router::new();
579 for (route, original) in icons {
580 let hashed: String = cache_buster.get_file(&original);
582 router = router.nest_service(route, ServeFile::new(hashed));
583 }
584 router
585}
586
587#[cfg(feature = "pages")]
593fn proxy_routes<S>(
594 frontend: &super::frontend::Frontend<S>,
595) -> (Router<WebServerState<S>>, Router<WebServerState<S>>)
596where
597 S: Clone + Send + Sync + 'static,
598{
599 use crate::analytics::{AnalyticsConfig, relay_envelope, relay_event, relay_script};
600 use axum::body::Bytes;
601 use axum::extract::ConnectInfo;
602 use axum::http::HeaderMap;
603 use axum::routing::post;
604
605 let client: reqwest::Client = reqwest::Client::new();
606 let paths: &crate::templates::FrontendRuntime = &frontend.runtime;
607
608 let analytics_script_upstream: String = frontend.analytics.upstream_script_url();
609 let analytics_event_upstream: String = AnalyticsConfig::upstream_event_url();
610 let sentry_script_upstream: String = frontend.sentry_dsn.upstream_script_url();
611 let sentry_envelope_upstream: String = frontend.sentry_dsn.upstream_envelope_url();
612
613 let scripts: Router<WebServerState<S>> = Router::new()
614 .route(
615 &paths.analytics.script_path,
616 get({
617 let client: reqwest::Client = client.clone();
618 let upstream: std::sync::Arc<str> =
619 std::sync::Arc::from(analytics_script_upstream.as_str());
620 move || {
621 let client: reqwest::Client = client.clone();
622 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
623 async move { relay_script(&client, &upstream).await }
624 }
625 }),
626 )
627 .route(
628 &paths.sentry_browser.script_path,
629 get({
630 let client: reqwest::Client = client.clone();
631 let upstream: std::sync::Arc<str> =
632 std::sync::Arc::from(sentry_script_upstream.as_str());
633 move || {
634 let client: reqwest::Client = client.clone();
635 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
636 async move { relay_script(&client, &upstream).await }
637 }
638 }),
639 );
640
641 let event_path: String = strip_api_prefix(&paths.analytics.event_path);
643 let tunnel_path: String = strip_api_prefix(&paths.sentry_browser.tunnel_path);
644
645 let endpoints: Router<WebServerState<S>> = Router::new()
646 .route(
647 &event_path,
648 post({
649 let client: reqwest::Client = client.clone();
650 let upstream: std::sync::Arc<str> =
651 std::sync::Arc::from(analytics_event_upstream.as_str());
652 move |ConnectInfo(peer): ConnectInfo<SocketAddr>,
653 headers: HeaderMap,
654 body: Bytes| {
655 let client: reqwest::Client = client.clone();
656 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
657 async move { relay_event(&client, &upstream, &headers, peer, body).await }
658 }
659 }),
660 )
661 .route(
662 &tunnel_path,
663 post({
664 let upstream: std::sync::Arc<str> =
665 std::sync::Arc::from(sentry_envelope_upstream.as_str());
666 move |body: Bytes| {
667 let client: reqwest::Client = client.clone();
668 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
669 async move { relay_envelope(&client, &upstream, body).await }
670 }
671 }),
672 );
673
674 (scripts, endpoints)
675}
676
677#[cfg(feature = "pages")]
680fn strip_api_prefix(path: &str) -> String {
681 path.strip_prefix(API_PREFIX)
682 .map_or_else(|| path.to_string(), String::from)
683}
684
685#[cfg(test)]
686mod tests {
687 use std::sync::Arc;
688 use std::sync::atomic::{AtomicBool, Ordering};
689 use std::time::Duration;
690
691 use axum::Router;
692 use axum::routing::get;
693 use tokio::time::timeout;
694
695 use super::super::shutdown::Shutdown;
696 use super::{
697 API_PREFIX, AppShutdown, DEFAULT_BODY_LIMIT, DEFAULT_DRAIN_TIMEOUT, WebServerError, health,
698 serve_on,
699 };
700
701 #[derive(Clone)]
704 struct RecordingState {
705 ran: Arc<AtomicBool>,
706 linger: Option<Duration>,
707 }
708
709 impl RecordingState {
710 fn instant() -> Self {
711 Self {
712 ran: Arc::new(AtomicBool::new(false)),
713 linger: None,
714 }
715 }
716 }
717
718 impl AppShutdown for RecordingState {
719 async fn on_shutdown(&self) {
720 if let Some(linger) = self.linger {
721 tokio::time::sleep(linger).await;
722 }
723 self.ran.store(true, Ordering::SeqCst);
724 }
725 }
726
727 async fn serve_until_shutdown(
729 state: RecordingState,
730 shutdown: Shutdown,
731 ) -> Result<(), WebServerError> {
732 let router: Router = Router::new().route("/health", get(health));
733 let cleanup: RecordingState = state.clone();
734 serve_on(router, "127.0.0.1", 0, shutdown, async move {
736 cleanup.on_shutdown().await;
737 })
738 .await
739 }
740
741 #[tokio::test]
742 async fn app_cleanup_runs_when_the_shutdown_signal_arrives() {
743 let shutdown: Shutdown = Shutdown::manual();
744 let state: RecordingState = RecordingState::instant();
745 let ran: Arc<AtomicBool> = Arc::clone(&state.ran);
746
747 shutdown.trigger();
748 timeout(
749 Duration::from_secs(1),
750 serve_until_shutdown(state, shutdown.clone()),
751 )
752 .await
753 .expect("serving ended promptly")
754 .expect("serving ended cleanly");
755
756 let expected: bool = true;
757 let actual: bool = ran.load(Ordering::SeqCst);
758 assert_eq!(expected, actual);
759 }
760
761 #[tokio::test]
762 async fn app_cleanup_never_runs_when_the_server_dies_before_any_signal() {
763 let shutdown: Shutdown = Shutdown::manual();
767 let state: RecordingState = RecordingState::instant();
768 let ran: Arc<AtomicBool> = Arc::clone(&state.ran);
769 let cleanup: RecordingState = state.clone();
770
771 let router: Router = Router::new().route("/health", get(health));
772 let result: Result<(), WebServerError> = timeout(
773 Duration::from_secs(1),
774 serve_on(router, "not-an-ip", 0, shutdown, async move {
775 cleanup.on_shutdown().await;
776 }),
777 )
778 .await
779 .expect("a bind failure returns immediately, it does not wait for cleanup");
780
781 assert!(result.is_err(), "an unparseable host is a bind error");
782
783 let expected: bool = false;
784 let actual: bool = ran.load(Ordering::SeqCst);
785 assert_eq!(expected, actual);
786 }
787
788 #[tokio::test(start_paused = true)]
789 async fn a_cleanup_that_overruns_the_window_is_cancelled_rather_than_awaited() {
790 let shutdown: Shutdown = Shutdown::manual();
794 let state: RecordingState = RecordingState {
795 ran: Arc::new(AtomicBool::new(false)),
796 linger: Some(DEFAULT_DRAIN_TIMEOUT * 2),
797 };
798 let ran: Arc<AtomicBool> = Arc::clone(&state.ran);
799
800 shutdown.trigger();
801 serve_until_shutdown(state, shutdown.clone())
802 .await
803 .expect("serving still ends cleanly when cleanup is cancelled");
804
805 let expected: bool = false;
806 let actual: bool = ran.load(Ordering::SeqCst);
807 assert_eq!(
808 expected, actual,
809 "the hook was cancelled at its await, so it never reached its final store"
810 );
811 }
812
813 #[tokio::test]
814 async fn a_server_with_nothing_in_flight_stops_at_once_instead_of_waiting_out_the_window() {
815 let shutdown: Shutdown = Shutdown::manual();
816 let router: Router = Router::new().route("/health", get(health));
817
818 let serving: tokio::task::JoinHandle<Result<(), WebServerError>> = tokio::spawn({
820 let shutdown: Shutdown = shutdown.clone();
821 async move { serve_on(router, "127.0.0.1", 0, shutdown, async {}).await }
822 });
823
824 shutdown.trigger();
825
826 timeout(Duration::from_secs(1), serving)
829 .await
830 .expect("the server returned as soon as the drain began")
831 .expect("the serving task did not panic")
832 .expect("serving ended cleanly");
833 }
834
835 #[test]
836 fn the_body_limit_accommodates_an_ordinary_form_post() {
837 let expected: usize = 256 * 1024;
840 let actual: usize = DEFAULT_BODY_LIMIT;
841 assert_eq!(expected, actual);
842 }
843
844 #[cfg(feature = "pages")]
845 #[test]
846 fn stripping_the_prefix_leaves_a_nestable_path() {
847 let expected: String = String::from("/boggledygook-a3f2c1d8");
848 let actual: String = super::strip_api_prefix("/api/v1/boggledygook-a3f2c1d8");
849 assert_eq!(expected, actual);
850 }
851
852 #[test]
853 fn the_api_prefix_is_the_one_every_project_shares() {
854 assert_eq!("/api/v1", API_PREFIX);
855 }
856}