1use std::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, info, instrument, warn};
16
17use crate::env;
18use crate::environment::Environment;
19
20use super::error::WebServerError;
21use super::shutdown::{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 let Self {
193 host,
194 port,
195 environment,
196 app,
197 body_limit,
198 router,
199 #[cfg(feature = "pages")]
200 frontend,
201 } = self;
202
203 let cache_buster: crate::assets::CacheBuster = crate::assets::CacheBuster::load()?;
206
207 let mut no_cache: Router<WebServerState<S>> = router;
208 let mut built_in: Router<WebServerState<S>> = Router::new().route("/health", get(health));
209
210 #[cfg(feature = "pages")]
211 let mut frontend_runtime: Option<crate::templates::FrontendRuntime> = None;
212 #[cfg(feature = "pages")]
213 let mut base: Option<crate::templates::BaseTemplateData> = None;
214 #[cfg(feature = "pages")]
215 let mut templates: Option<crate::templates::TemplateRegistry<'static>> = None;
216 #[cfg(feature = "pages")]
217 let mut not_found = None;
218 #[cfg(feature = "pages")]
219 let mut proxy_scripts: Option<Router<WebServerState<S>>> = None;
220
221 #[cfg(feature = "pages")]
222 if let Some(params) = frontend {
223 let registry: crate::templates::TemplateRegistry<'static> =
224 crate::templates::TemplateRegistry::from_dir(crate::templates::TEMPLATE_ROOT)?;
225
226 let built = super::frontend::Frontend::build(params, &cache_buster, environment)?;
227
228 no_cache = no_cache.merge(well_known_routes(&built.well_known));
229 no_cache = no_cache.merge(icon_routes(&cache_buster, built.has_svg_icon));
230
231 let (scripts, endpoints) = proxy_routes(&built);
232 proxy_scripts = Some(scripts);
236 built_in = built_in.merge(endpoints);
237
238 not_found = Some(built.not_found.clone());
239 no_cache = no_cache.merge(built.pages.into_router());
240
241 frontend_runtime = Some(built.runtime);
242 base = Some(built.base);
243 templates = Some(registry);
244 }
245
246 let no_cache: Router<WebServerState<S>> = no_cache.nest(API_PREFIX, built_in);
247
248 let mut app_router: Router<WebServerState<S>> = apply_cache_policy(no_cache, &cache_buster);
249
250 #[cfg(feature = "pages")]
254 if let Some(scripts) = proxy_scripts {
255 app_router = app_router.merge(scripts);
256 }
257
258 #[cfg(feature = "pages")]
259 let app_router = match not_found {
260 Some((page, data)) => app_router.fallback(move |axum::extract::State(state)| {
261 let page = std::sync::Arc::clone(&page);
262 let data = std::sync::Arc::clone(&data);
263 async move {
264 let body: Response = super::pages::render_or_500(&state, &page, &data);
265 (StatusCode::NOT_FOUND, body).into_response()
266 }
267 }),
268 None => app_router.fallback(plain_not_found),
269 };
270 #[cfg(not(feature = "pages"))]
271 let app_router = app_router.fallback(plain_not_found);
272
273 let state: WebServerState<S> = WebServerState::new(StateParts {
274 host: host.clone(),
275 port,
276 environment,
277 shutdown: shutdown.clone(),
278 #[cfg(feature = "templates")]
279 base,
280 #[cfg(feature = "templates")]
281 templates,
282 cache_buster: Some(cache_buster),
283 #[cfg(feature = "templates")]
284 frontend: frontend_runtime,
285 app,
286 });
287
288 let app_router = app_router
291 .with_state(state)
292 .layer(
293 TraceLayer::new_for_http()
294 .make_span_with(
300 DefaultMakeSpan::new()
301 .level(Level::INFO)
302 .include_headers(false),
303 )
304 .on_response(
305 DefaultOnResponse::new()
306 .level(Level::INFO)
307 .latency_unit(LatencyUnit::Millis),
308 ),
309 )
310 .layer(DefaultBodyLimit::max(body_limit));
311
312 serve_on(app_router, &host, port, shutdown).await
313 }
314}
315
316async fn serve_on(
319 router: Router,
320 host: &str,
321 port: u16,
322 shutdown: Shutdown,
323) -> Result<(), WebServerError> {
324 let ip: IpAddr = IpAddr::from_str(host).map_err(|source| WebServerError::Bind {
325 addr: SocketAddr::from(([0, 0, 0, 0], port)),
326 source: std::io::Error::new(std::io::ErrorKind::InvalidInput, source),
327 })?;
328 let address: SocketAddr = SocketAddr::new(ip, port);
329 let listener: TcpListener =
330 TcpListener::bind(address)
331 .await
332 .map_err(|source| WebServerError::Bind {
333 addr: address,
334 source,
335 })?;
336
337 info!("listening on http://{address}");
338
339 let serving = serve(
340 listener,
341 router.into_make_service_with_connect_info::<SocketAddr>(),
342 )
343 .with_graceful_shutdown(shutdown.clone().recv())
344 .into_future();
345 tokio::pin!(serving);
346
347 tokio::select! {
351 result = &mut serving => return result.map_err(WebServerError::Serve),
352 () = shutdown.recv() => {}
353 }
354
355 match tokio::time::timeout(DEFAULT_DRAIN_TIMEOUT, &mut serving).await {
356 Ok(result) => result.map_err(WebServerError::Serve),
357 Err(_elapsed) => {
360 warn!("drain window elapsed after {DEFAULT_DRAIN_TIMEOUT:?}; closing what remained");
361 Ok(())
362 }
363 }
364}
365
366fn apply_cache_policy<S>(
369 router: Router<WebServerState<S>>,
370 cache_buster: &crate::assets::CacheBuster,
371) -> Router<WebServerState<S>>
372where
373 S: Clone + Send + Sync + 'static,
374{
375 let router: Router<WebServerState<S>> = router.layer(axum::middleware::from_fn(
376 crate::assets::CacheBuster::never_cache_middleware,
377 ));
378
379 if cache_buster.is_empty() {
380 return router;
381 }
382
383 router.merge(
384 Router::new()
385 .nest_service(
386 "/static",
387 tower_http::services::ServeDir::new(crate::assets::STATIC_DIRECTORY),
388 )
389 .layer(axum::middleware::from_fn(
390 crate::assets::CacheBuster::forever_cache_middleware,
391 )),
392 )
393}
394
395async fn health() -> StatusCode {
398 StatusCode::OK
399}
400
401async fn plain_not_found() -> Response {
403 (StatusCode::NOT_FOUND, "404").into_response()
404}
405
406#[cfg(feature = "pages")]
413fn well_known_routes<S>(well_known: &super::frontend::WellKnown) -> Router<WebServerState<S>>
414where
415 S: Clone + Send + Sync + 'static,
416{
417 use axum::http::header;
418
419 fn text<S>(
420 router: Router<WebServerState<S>>,
421 path: &str,
422 content_type: &'static str,
423 body: String,
424 ) -> Router<WebServerState<S>>
425 where
426 S: Clone + Send + Sync + 'static,
427 {
428 router.route(
429 path,
430 get(move || {
431 let body: String = body.clone();
432 async move { ([(header::CONTENT_TYPE, content_type)], body) }
433 }),
434 )
435 }
436
437 let mut router: Router<WebServerState<S>> = Router::new();
438 router = text(
439 router,
440 "/robots.txt",
441 "text/plain; charset=utf-8",
442 well_known.robots_txt.clone(),
443 );
444 router = text(
445 router,
446 "/humans.txt",
447 "text/plain; charset=utf-8",
448 well_known.humans_txt.clone(),
449 );
450 router = text(
451 router,
452 "/site.webmanifest",
453 "application/manifest+json",
454 well_known.webmanifest.clone(),
455 );
456 router = text(
457 router,
458 crate::sitemap::SITEMAP_INDEX_PATH,
459 "application/xml",
460 well_known.sitemaps.index().to_string(),
461 );
462 for (index, chunk) in well_known.sitemaps.chunks().iter().enumerate() {
463 router = text(
464 router,
465 &format!("/sitemap-{}.xml", index + 1),
466 "application/xml",
467 chunk.clone(),
468 );
469 }
470 router
471}
472
473#[cfg(feature = "pages")]
479fn icon_routes<S>(
480 cache_buster: &crate::assets::CacheBuster,
481 has_svg_icon: bool,
482) -> Router<WebServerState<S>>
483where
484 S: Clone + Send + Sync + 'static,
485{
486 use tower_http::services::ServeFile;
487
488 let mut icons: Vec<(&str, String)> = vec![
489 (
490 "/favicon.ico",
491 String::from("static/image/favicon/favicon.ico"),
492 ),
493 (
494 "/apple-touch-icon.png",
495 String::from("static/image/favicon/apple-touch-icon.png"),
496 ),
497 (
498 "/icon-192.png",
499 String::from("static/image/favicon/icon-192.png"),
500 ),
501 (
502 "/icon-512.png",
503 String::from("static/image/favicon/icon-512.png"),
504 ),
505 ];
506 if has_svg_icon {
508 icons.push((
509 "/favicon.svg",
510 String::from("static/image/favicon/favicon.svg"),
511 ));
512 }
513
514 let mut router: Router<WebServerState<S>> = Router::new();
515 for (route, original) in icons {
516 let hashed: String = cache_buster.get_file(&original);
518 router = router.nest_service(route, ServeFile::new(hashed));
519 }
520 router
521}
522
523#[cfg(feature = "pages")]
529fn proxy_routes<S>(
530 frontend: &super::frontend::Frontend<S>,
531) -> (Router<WebServerState<S>>, Router<WebServerState<S>>)
532where
533 S: Clone + Send + Sync + 'static,
534{
535 use crate::analytics::{AnalyticsConfig, relay_envelope, relay_event, relay_script};
536 use axum::body::Bytes;
537 use axum::extract::ConnectInfo;
538 use axum::http::HeaderMap;
539 use axum::routing::post;
540
541 let client: reqwest::Client = reqwest::Client::new();
542 let paths = &frontend.runtime;
543
544 let analytics_script_upstream: String = frontend.analytics.upstream_script_url();
545 let analytics_event_upstream: String = AnalyticsConfig::upstream_event_url();
546 let sentry_script_upstream: String = frontend.sentry_dsn.upstream_script_url();
547 let sentry_envelope_upstream: String = frontend.sentry_dsn.upstream_envelope_url();
548
549 let scripts: Router<WebServerState<S>> = Router::new()
550 .route(
551 &paths.analytics.script_path,
552 get({
553 let client: reqwest::Client = client.clone();
554 let upstream: std::sync::Arc<str> =
555 std::sync::Arc::from(analytics_script_upstream.as_str());
556 move || {
557 let client: reqwest::Client = client.clone();
558 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
559 async move { relay_script(&client, &upstream).await }
560 }
561 }),
562 )
563 .route(
564 &paths.sentry_browser.script_path,
565 get({
566 let client: reqwest::Client = client.clone();
567 let upstream: std::sync::Arc<str> =
568 std::sync::Arc::from(sentry_script_upstream.as_str());
569 move || {
570 let client: reqwest::Client = client.clone();
571 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
572 async move { relay_script(&client, &upstream).await }
573 }
574 }),
575 );
576
577 let event_path: String = strip_api_prefix(&paths.analytics.event_path);
579 let tunnel_path: String = strip_api_prefix(&paths.sentry_browser.tunnel_path);
580
581 let endpoints: Router<WebServerState<S>> = Router::new()
582 .route(
583 &event_path,
584 post({
585 let client: reqwest::Client = client.clone();
586 let upstream: std::sync::Arc<str> =
587 std::sync::Arc::from(analytics_event_upstream.as_str());
588 move |ConnectInfo(peer): ConnectInfo<SocketAddr>,
589 headers: HeaderMap,
590 body: Bytes| {
591 let client: reqwest::Client = client.clone();
592 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
593 async move { relay_event(&client, &upstream, &headers, peer, body).await }
594 }
595 }),
596 )
597 .route(
598 &tunnel_path,
599 post({
600 let upstream: std::sync::Arc<str> =
601 std::sync::Arc::from(sentry_envelope_upstream.as_str());
602 move |body: Bytes| {
603 let client: reqwest::Client = client.clone();
604 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
605 async move { relay_envelope(&client, &upstream, body).await }
606 }
607 }),
608 );
609
610 (scripts, endpoints)
611}
612
613#[cfg(feature = "pages")]
616fn strip_api_prefix(path: &str) -> String {
617 path.strip_prefix(API_PREFIX)
618 .map_or_else(|| path.to_string(), String::from)
619}
620
621#[cfg(test)]
622mod tests {
623 use std::time::Duration;
624
625 use axum::Router;
626 use axum::routing::get;
627 use tokio::time::timeout;
628
629 use super::super::shutdown::Shutdown;
630 use super::{API_PREFIX, DEFAULT_BODY_LIMIT, health, serve_on};
631
632 #[tokio::test]
633 async fn a_server_with_nothing_in_flight_stops_at_once_instead_of_waiting_out_the_window() {
634 let shutdown: Shutdown = Shutdown::manual();
635 let router: Router = Router::new().route("/health", get(health));
636
637 let serving = tokio::spawn({
639 let shutdown: Shutdown = shutdown.clone();
640 async move { serve_on(router, "127.0.0.1", 0, shutdown).await }
641 });
642
643 shutdown.trigger();
644
645 timeout(Duration::from_secs(1), serving)
648 .await
649 .expect("the server returned as soon as the drain began")
650 .expect("the serving task did not panic")
651 .expect("serving ended cleanly");
652 }
653
654 #[test]
655 fn the_body_limit_accommodates_an_ordinary_form_post() {
656 let expected: usize = 256 * 1024;
659 let actual: usize = DEFAULT_BODY_LIMIT;
660 assert_eq!(expected, actual);
661 }
662
663 #[cfg(feature = "pages")]
664 #[test]
665 fn stripping_the_prefix_leaves_a_nestable_path() {
666 let expected: String = String::from("/boggledygook-a3f2c1d8");
667 let actual: String = super::strip_api_prefix("/api/v1/boggledygook-a3f2c1d8");
668 assert_eq!(expected, actual);
669 }
670
671 #[test]
672 fn the_api_prefix_is_the_one_every_project_shares() {
673 assert_eq!("/api/v1", API_PREFIX);
674 }
675}