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 #[cfg(feature = "pages")]
67 feed: Option<crate::feed::Feed>,
68}
69
70impl WebServer<()> {
71 #[must_use]
73 pub fn new(host: impl Into<String>, port: u16, environment: Environment) -> Self {
74 Self::with_state(host, port, environment, ())
75 }
76
77 pub fn from_env() -> Result<Self, WebServerError> {
87 Self::from_env_with_state(())
88 }
89}
90
91impl<S> WebServer<S>
92where
93 S: Clone + Send + Sync + 'static,
94{
95 #[must_use]
97 pub fn with_state(
98 host: impl Into<String>,
99 port: u16,
100 environment: Environment,
101 app: S,
102 ) -> Self {
103 Self {
104 host: host.into(),
105 port,
106 environment,
107 app,
108 body_limit: DEFAULT_BODY_LIMIT,
109 router: Router::new(),
110 #[cfg(feature = "pages")]
111 frontend: None,
112 #[cfg(feature = "pages")]
113 feed: None,
114 }
115 }
116
117 pub fn from_env_with_state(app: S) -> Result<Self, WebServerError> {
124 let environment: Environment = Environment::from_env()?;
125 let host: String = env::optional(ENV_HOST).unwrap_or_else(|| {
126 if environment.is_production() {
127 String::from("0.0.0.0")
128 } else {
129 String::from("127.0.0.1")
130 }
131 });
132 let port: u16 =
133 env::parse_or(ENV_PORT, "port number", DEFAULT_PORT).map_err(WebServerError::Env)?;
134
135 Ok(Self::with_state(host, port, environment, app))
136 }
137
138 #[must_use]
143 pub const fn body_limit(mut self, body_limit: usize) -> Self {
144 self.body_limit = body_limit;
145 self
146 }
147
148 #[must_use]
150 pub fn nest(mut self, path: &str, router: Router<WebServerState<S>>) -> Self {
151 self.router = self.router.nest(path, router);
152 self
153 }
154
155 #[must_use]
157 pub fn merge(mut self, router: Router<WebServerState<S>>) -> Self {
158 self.router = self.router.merge(router);
159 self
160 }
161
162 #[must_use]
164 pub fn nest_service<T>(mut self, path: &str, service: T) -> Self
165 where
166 T: tower::Service<axum::extract::Request, Error = std::convert::Infallible>
167 + Clone
168 + Send
169 + Sync
170 + 'static,
171 T::Response: axum::response::IntoResponse,
172 T::Future: Send + 'static,
173 {
174 self.router = self.router.nest_service(path, service);
175 self
176 }
177
178 #[cfg(feature = "pages")]
185 #[must_use]
186 pub fn frontend(mut self, params: super::frontend::FrontendParams<S>) -> Self {
187 self.frontend = Some(params);
188 self
189 }
190
191 #[cfg(feature = "pages")]
204 #[must_use]
205 pub fn feed(mut self, feed: crate::feed::Feed) -> Self {
206 self.feed = Some(feed);
207 self
208 }
209
210 #[instrument(skip_all)]
217 pub async fn run(self, shutdown: Shutdown) -> Result<(), WebServerError>
218 where
219 S: AppShutdown,
220 {
221 let Self {
222 host,
223 port,
224 environment,
225 app,
226 body_limit,
227 router,
228 #[cfg(feature = "pages")]
229 frontend,
230 #[cfg(feature = "pages")]
231 feed,
232 } = self;
233
234 #[cfg(feature = "pages")]
237 if feed.is_some() && frontend.is_none() {
238 return Err(WebServerError::FeedWithoutFrontend);
239 }
240
241 let cache_buster: crate::assets::CacheBuster = crate::assets::CacheBuster::load()?;
244
245 let mut no_cache: Router<WebServerState<S>> = router;
246 let mut built_in: Router<WebServerState<S>> = Router::new().route("/health", get(health));
247
248 #[cfg(feature = "pages")]
249 let mut frontend_runtime: Option<crate::templates::FrontendRuntime> = None;
250 #[cfg(feature = "pages")]
251 let mut base: Option<crate::templates::BaseTemplateData> = None;
252 #[cfg(feature = "pages")]
253 let mut templates: Option<crate::templates::TemplateRegistry<'static>> = None;
254 #[cfg(feature = "pages")]
255 let mut not_found: Option<(
256 std::sync::Arc<crate::templates::PageTemplateData>,
257 std::sync::Arc<serde_json::Value>,
258 )> = None;
259 #[cfg(feature = "pages")]
260 let mut proxy_scripts: Option<Router<WebServerState<S>>> = None;
261 #[cfg(feature = "pages")]
262 let mut feed_router: Option<Router<WebServerState<S>>> = None;
263
264 log_immutable_assets(&cache_buster);
265
266 #[cfg(feature = "pages")]
267 if let Some(params) = frontend {
268 let assembled: Assembled<S> =
269 assemble_frontend(params, feed, &cache_buster, environment)?;
270
271 no_cache = no_cache.merge(assembled.routes);
272 built_in = built_in.merge(assembled.api);
273 proxy_scripts = Some(assembled.proxy_scripts);
274 feed_router = assembled.feeds;
275 not_found = Some(assembled.not_found);
276 frontend_runtime = Some(assembled.runtime);
277 base = Some(assembled.base);
278 templates = Some(assembled.templates);
279 }
280
281 let no_cache: Router<WebServerState<S>> = no_cache.nest(API_PREFIX, built_in);
282
283 let mut app_router: Router<WebServerState<S>> = apply_cache_policy(no_cache, &cache_buster);
284
285 #[cfg(feature = "pages")]
289 if let Some(scripts) = proxy_scripts {
290 app_router = app_router.merge(scripts);
291 }
292
293 #[cfg(feature = "pages")]
294 if let Some(feeds) = feed_router {
295 app_router = app_router.merge(feeds);
296 }
297
298 #[cfg(feature = "pages")]
299 let app_router: Router<WebServerState<S>> = attach_not_found(app_router, not_found);
300 #[cfg(not(feature = "pages"))]
301 let app_router: Router<WebServerState<S>> = app_router.fallback(plain_not_found);
302
303 let state: WebServerState<S> = WebServerState::new(StateParts {
304 host: host.clone(),
305 port,
306 environment,
307 shutdown: shutdown.clone(),
308 #[cfg(feature = "templates")]
309 base,
310 #[cfg(feature = "templates")]
311 templates,
312 cache_buster: Some(cache_buster),
313 #[cfg(feature = "templates")]
314 frontend: frontend_runtime,
315 app,
316 });
317
318 let cleanup: WebServerState<S> = state.clone();
321
322 let app_router: Router = app_router
325 .with_state(state)
326 .layer(
327 TraceLayer::new_for_http()
328 .make_span_with(
334 DefaultMakeSpan::new()
335 .level(Level::INFO)
336 .include_headers(false),
337 )
338 .on_response(
339 DefaultOnResponse::new()
340 .level(Level::INFO)
341 .latency_unit(LatencyUnit::Millis),
342 ),
343 )
344 .layer(DefaultBodyLimit::max(body_limit));
345
346 serve_on(app_router, &host, port, shutdown, async move {
347 cleanup.app().on_shutdown().await;
348 })
349 .await
350 }
351}
352
353#[cfg(feature = "pages")]
360fn attach_not_found<S>(
361 router: Router<WebServerState<S>>,
362 not_found: Option<(
363 std::sync::Arc<crate::templates::PageTemplateData>,
364 std::sync::Arc<serde_json::Value>,
365 )>,
366) -> Router<WebServerState<S>>
367where
368 S: Clone + Send + Sync + 'static,
369{
370 match not_found {
371 Some((page, data)) => router.fallback(move |axum::extract::State(state)| {
372 let page: std::sync::Arc<crate::templates::PageTemplateData> =
373 std::sync::Arc::clone(&page);
374 let data: std::sync::Arc<serde_json::Value> = std::sync::Arc::clone(&data);
375 async move {
376 let body: Response = super::pages::render_or_500(&state, &page, &data);
377 (StatusCode::NOT_FOUND, body).into_response()
378 }
379 }),
380 None => router.fallback(plain_not_found),
381 }
382}
383
384async fn serve_on<F>(
391 router: Router,
392 host: &str,
393 port: u16,
394 shutdown: Shutdown,
395 on_shutdown: F,
396) -> Result<(), WebServerError>
397where
398 F: Future<Output = ()> + Send,
399{
400 let ip: IpAddr = IpAddr::from_str(host).map_err(|source| WebServerError::Bind {
401 addr: SocketAddr::from(([0, 0, 0, 0], port)),
402 source: std::io::Error::new(std::io::ErrorKind::InvalidInput, source),
403 })?;
404 let address: SocketAddr = SocketAddr::new(ip, port);
405 let listener: TcpListener =
406 TcpListener::bind(address)
407 .await
408 .map_err(|source| WebServerError::Bind {
409 addr: address,
410 source,
411 })?;
412
413 info!("listening on http://{address}");
414
415 let serving = serve(
419 listener,
420 router.into_make_service_with_connect_info::<SocketAddr>(),
421 )
422 .with_graceful_shutdown(shutdown.clone().recv())
423 .into_future();
424 tokio::pin!(serving);
425
426 tokio::select! {
430 result = &mut serving => return result.map_err(WebServerError::Serve),
431 () = shutdown.recv() => {}
432 }
433
434 let (served, cleaned) = tokio::join!(
439 tokio::time::timeout(DEFAULT_DRAIN_TIMEOUT, &mut serving),
440 tokio::time::timeout(DEFAULT_DRAIN_TIMEOUT, on_shutdown),
441 );
442
443 if cleaned.is_err() {
447 error!(
448 "app shutdown hook exceeded {DEFAULT_DRAIN_TIMEOUT:?}; cleanup was cancelled part-way"
449 );
450 }
451
452 match served {
453 Ok(result) => result.map_err(WebServerError::Serve),
454 Err(_elapsed) => {
457 warn!("drain window elapsed after {DEFAULT_DRAIN_TIMEOUT:?}; closing what remained");
458 Ok(())
459 }
460 }
461}
462
463fn apply_cache_policy<S>(
466 router: Router<WebServerState<S>>,
467 cache_buster: &crate::assets::CacheBuster,
468) -> Router<WebServerState<S>>
469where
470 S: Clone + Send + Sync + 'static,
471{
472 let router: Router<WebServerState<S>> = router.layer(axum::middleware::from_fn(
473 crate::assets::CacheBuster::never_cache_middleware,
474 ));
475
476 if cache_buster.is_empty() {
477 return router;
478 }
479
480 router.merge(
481 Router::new()
482 .nest_service(
483 "/static",
484 tower_http::services::ServeDir::new(crate::assets::STATIC_DIRECTORY),
485 )
486 .layer(axum::middleware::from_fn(
487 crate::assets::CacheBuster::forever_cache_middleware,
488 )),
489 )
490}
491
492async fn health() -> StatusCode {
495 StatusCode::OK
496}
497
498async fn plain_not_found() -> Response {
500 (StatusCode::NOT_FOUND, "404").into_response()
501}
502
503#[cfg(feature = "pages")]
510fn well_known_routes<S>(well_known: &super::frontend::WellKnown) -> Router<WebServerState<S>>
511where
512 S: Clone + Send + Sync + 'static,
513{
514 use axum::http::header;
515
516 fn text<S>(
517 router: Router<WebServerState<S>>,
518 path: &str,
519 content_type: &'static str,
520 body: String,
521 ) -> Router<WebServerState<S>>
522 where
523 S: Clone + Send + Sync + 'static,
524 {
525 router.route(
526 path,
527 get(move || {
528 let body: String = body.clone();
529 async move { ([(header::CONTENT_TYPE, content_type)], body) }
530 }),
531 )
532 }
533
534 let mut router: Router<WebServerState<S>> = Router::new();
535 router = text(
536 router,
537 "/robots.txt",
538 "text/plain; charset=utf-8",
539 well_known.robots_txt.clone(),
540 );
541 router = text(
542 router,
543 "/humans.txt",
544 "text/plain; charset=utf-8",
545 well_known.humans_txt.clone(),
546 );
547 router = text(
548 router,
549 "/site.webmanifest",
550 "application/manifest+json",
551 well_known.webmanifest.clone(),
552 );
553 router = text(
554 router,
555 crate::sitemap::SITEMAP_INDEX_PATH,
556 "application/xml",
557 well_known.sitemaps.index().to_string(),
558 );
559 for (index, chunk) in well_known.sitemaps.chunks().iter().enumerate() {
560 router = text(
561 router,
562 &format!("/sitemap-{}.xml", index + 1),
563 "application/xml",
564 chunk.clone(),
565 );
566 }
567 router
568}
569
570#[cfg(feature = "pages")]
576fn icon_routes<S>(
577 cache_buster: &crate::assets::CacheBuster,
578 has_svg_icon: bool,
579) -> Router<WebServerState<S>>
580where
581 S: Clone + Send + Sync + 'static,
582{
583 use tower_http::services::ServeFile;
584
585 let mut icons: Vec<(&str, String)> = vec![
586 (
587 "/favicon.ico",
588 String::from("static/image/favicon/favicon.ico"),
589 ),
590 (
591 "/apple-touch-icon.png",
592 String::from("static/image/favicon/apple-touch-icon.png"),
593 ),
594 (
595 "/icon-192.png",
596 String::from("static/image/favicon/icon-192.png"),
597 ),
598 (
599 "/icon-512.png",
600 String::from("static/image/favicon/icon-512.png"),
601 ),
602 ];
603 if has_svg_icon {
605 icons.push((
606 "/favicon.svg",
607 String::from("static/image/favicon/favicon.svg"),
608 ));
609 }
610
611 let mut router: Router<WebServerState<S>> = Router::new();
612 for (route, original) in icons {
613 let hashed: String = cache_buster.get_file(&original);
615 router = router.nest_service(route, ServeFile::new(hashed));
616 }
617 router
618}
619
620#[cfg(feature = "pages")]
631fn feed_routes<S>(feeds: &crate::feed::FeedSet) -> Router<WebServerState<S>>
632where
633 S: Clone + Send + Sync + 'static,
634{
635 use axum::http::{HeaderMap, HeaderValue, StatusCode, header};
636
637 let last_modified: HeaderValue = HeaderValue::from_str(
641 &feeds
642 .last_modified()
643 .format("%a, %d %b %Y %H:%M:%S GMT")
644 .to_string(),
645 )
646 .unwrap_or_else(|_| HeaderValue::from_static(""));
647 let cache_control: HeaderValue = HeaderValue::from_str(&format!(
648 "public, max-age={}",
649 crate::feed::FEED_MAX_AGE_SECONDS
650 ))
651 .unwrap_or_else(|_| HeaderValue::from_static("public"));
652
653 let mut router: Router<WebServerState<S>> = Router::new();
654 for document in feeds.documents() {
655 let body: String = String::from(document.body());
656 let etag_text: String = String::from(document.etag());
657 let etag: HeaderValue =
658 HeaderValue::from_str(&etag_text).unwrap_or_else(|_| HeaderValue::from_static(""));
659 let content_type: HeaderValue = HeaderValue::from_static(document.content_type());
660 let last_modified: HeaderValue = last_modified.clone();
661 let cache_control: HeaderValue = cache_control.clone();
662
663 router = router.route(
664 document.path(),
665 get(move |headers: HeaderMap| {
666 let body: String = body.clone();
667 let etag_text: String = etag_text.clone();
668 let etag: HeaderValue = etag.clone();
669 let content_type: HeaderValue = content_type.clone();
670 let last_modified: HeaderValue = last_modified.clone();
671 let cache_control: HeaderValue = cache_control.clone();
672 async move {
673 let fresh: bool = if_none_match(&headers, &etag_text);
674
675 let mut response: Response = if fresh {
680 StatusCode::NOT_MODIFIED.into_response()
681 } else {
682 (StatusCode::OK, body).into_response()
683 };
684
685 let headers: &mut HeaderMap = response.headers_mut();
686 headers.insert(header::CACHE_CONTROL, cache_control);
687 headers.insert(header::ETAG, etag);
688 headers.insert(header::LAST_MODIFIED, last_modified);
689 if !fresh {
690 headers.insert(header::CONTENT_TYPE, content_type);
691 }
692
693 response
694 }
695 }),
696 );
697 }
698
699 router
700}
701
702#[cfg(feature = "pages")]
708fn if_none_match(headers: &axum::http::HeaderMap, etag: &str) -> bool {
709 let Some(header) = headers
710 .get(axum::http::header::IF_NONE_MATCH)
711 .and_then(|value| value.to_str().ok())
712 else {
713 return false;
714 };
715
716 header.split(',').any(|candidate| {
717 let candidate: &str = candidate.trim();
718 candidate == "*" || candidate.trim_start_matches("W/") == etag
719 })
720}
721
722fn log_immutable_assets(cache_buster: &crate::assets::CacheBuster) {
732 for (original, hashed) in cache_buster.cache() {
733 info!("immutable 1y /{original} -> /{hashed}");
734 }
735}
736
737#[cfg(feature = "pages")]
739fn log_served_documents(well_known: &super::frontend::WellKnown) {
740 for path in ["/robots.txt", "/humans.txt", "/site.webmanifest"] {
741 info!("uncached {path}");
742 }
743 for path in well_known.sitemaps.paths() {
744 info!("uncached {path}");
745 }
746
747 if let Some(feeds) = well_known.feeds.as_ref() {
748 for document in feeds.documents() {
749 info!(
750 "cached {}m {} (etag {})",
751 crate::feed::FEED_MAX_AGE_SECONDS / 60,
752 document.path(),
753 document.etag()
754 );
755 }
756 }
757}
758
759#[cfg(feature = "pages")]
765struct Assembled<S> {
766 routes: Router<WebServerState<S>>,
768 api: Router<WebServerState<S>>,
770 proxy_scripts: Router<WebServerState<S>>,
772 feeds: Option<Router<WebServerState<S>>>,
774 runtime: crate::templates::FrontendRuntime,
775 base: crate::templates::BaseTemplateData,
776 templates: crate::templates::TemplateRegistry<'static>,
777 not_found: (
778 std::sync::Arc<crate::templates::PageTemplateData>,
779 std::sync::Arc<serde_json::Value>,
780 ),
781}
782
783#[cfg(feature = "pages")]
785fn assemble_frontend<S>(
786 params: super::frontend::FrontendParams<S>,
787 feed: Option<crate::feed::Feed>,
788 cache_buster: &crate::assets::CacheBuster,
789 environment: Environment,
790) -> Result<Assembled<S>, WebServerError>
791where
792 S: Clone + Send + Sync + 'static,
793{
794 let templates: crate::templates::TemplateRegistry<'static> =
795 crate::templates::TemplateRegistry::from_dir(crate::templates::TEMPLATE_ROOT)?;
796
797 let built: super::frontend::Frontend<S> =
798 super::frontend::Frontend::build(params, feed, cache_buster, environment)?;
799
800 log_served_documents(&built.well_known);
801
802 let feeds: Option<Router<WebServerState<S>>> = built.well_known.feeds.as_ref().map(feed_routes);
806
807 let mut routes: Router<WebServerState<S>> = well_known_routes(&built.well_known);
808 routes = routes.merge(icon_routes(cache_buster, built.has_svg_icon));
809
810 let (proxy_scripts, api) = proxy_routes(&built);
814
815 let not_found = built.not_found.clone();
816 routes = routes.merge(built.pages.into_router());
817
818 Ok(Assembled {
819 routes,
820 api,
821 proxy_scripts,
822 feeds,
823 runtime: built.runtime,
824 base: built.base,
825 templates,
826 not_found,
827 })
828}
829
830#[cfg(feature = "pages")]
836fn proxy_routes<S>(
837 frontend: &super::frontend::Frontend<S>,
838) -> (Router<WebServerState<S>>, Router<WebServerState<S>>)
839where
840 S: Clone + Send + Sync + 'static,
841{
842 use crate::analytics::{AnalyticsConfig, relay_envelope, relay_event, relay_script};
843 use axum::body::Bytes;
844 use axum::extract::ConnectInfo;
845 use axum::http::HeaderMap;
846 use axum::routing::post;
847
848 let client: reqwest::Client = reqwest::Client::new();
849 let paths: &crate::templates::FrontendRuntime = &frontend.runtime;
850
851 let analytics_script_upstream: String = frontend.analytics.upstream_script_url();
852 let analytics_event_upstream: String = AnalyticsConfig::upstream_event_url();
853 let sentry_script_upstream: String = frontend.sentry_dsn.upstream_script_url();
854 let sentry_envelope_upstream: String = frontend.sentry_dsn.upstream_envelope_url();
855
856 let scripts: Router<WebServerState<S>> = Router::new()
857 .route(
858 &paths.analytics.script_path,
859 get({
860 let client: reqwest::Client = client.clone();
861 let upstream: std::sync::Arc<str> =
862 std::sync::Arc::from(analytics_script_upstream.as_str());
863 move || {
864 let client: reqwest::Client = client.clone();
865 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
866 async move { relay_script(&client, &upstream).await }
867 }
868 }),
869 )
870 .route(
871 &paths.sentry_browser.script_path,
872 get({
873 let client: reqwest::Client = client.clone();
874 let upstream: std::sync::Arc<str> =
875 std::sync::Arc::from(sentry_script_upstream.as_str());
876 move || {
877 let client: reqwest::Client = client.clone();
878 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
879 async move { relay_script(&client, &upstream).await }
880 }
881 }),
882 );
883
884 let event_path: String = strip_api_prefix(&paths.analytics.event_path);
886 let tunnel_path: String = strip_api_prefix(&paths.sentry_browser.tunnel_path);
887
888 let endpoints: Router<WebServerState<S>> = Router::new()
889 .route(
890 &event_path,
891 post({
892 let client: reqwest::Client = client.clone();
893 let upstream: std::sync::Arc<str> =
894 std::sync::Arc::from(analytics_event_upstream.as_str());
895 move |ConnectInfo(peer): ConnectInfo<SocketAddr>,
896 headers: HeaderMap,
897 body: Bytes| {
898 let client: reqwest::Client = client.clone();
899 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
900 async move { relay_event(&client, &upstream, &headers, peer, body).await }
901 }
902 }),
903 )
904 .route(
905 &tunnel_path,
906 post({
907 let upstream: std::sync::Arc<str> =
908 std::sync::Arc::from(sentry_envelope_upstream.as_str());
909 move |body: Bytes| {
910 let client: reqwest::Client = client.clone();
911 let upstream: std::sync::Arc<str> = std::sync::Arc::clone(&upstream);
912 async move { relay_envelope(&client, &upstream, body).await }
913 }
914 }),
915 );
916
917 (scripts, endpoints)
918}
919
920#[cfg(feature = "pages")]
923fn strip_api_prefix(path: &str) -> String {
924 path.strip_prefix(API_PREFIX)
925 .map_or_else(|| path.to_string(), String::from)
926}
927
928#[cfg(test)]
929mod tests {
930
931 #[cfg(feature = "pages")]
932 mod conditional_get {
933 use axum::http::{HeaderMap, HeaderValue, header};
934
935 use crate::webserver::server::if_none_match;
936
937 const ETAG: &str = "\"abc123\"";
938
939 fn headers(value: &str) -> HeaderMap {
940 let mut headers: HeaderMap = HeaderMap::new();
941 headers.insert(
942 header::IF_NONE_MATCH,
943 HeaderValue::from_str(value).expect("a valid header"),
944 );
945 headers
946 }
947
948 #[test]
949 fn a_request_without_the_header_always_gets_the_body() {
950 let expected: bool = false;
951 let actual: bool = if_none_match(&HeaderMap::new(), ETAG);
952 assert_eq!(expected, actual);
953 }
954
955 #[test]
956 fn the_same_validator_means_the_reader_already_has_this_feed() {
957 let expected: bool = true;
958 let actual: bool = if_none_match(&headers(ETAG), ETAG);
959 assert_eq!(expected, actual);
960 }
961
962 #[test]
963 fn a_weak_validator_still_matches_because_feeds_need_no_byte_equality() {
964 let expected: bool = true;
965 let actual: bool = if_none_match(&headers("W/\"abc123\""), ETAG);
966 assert_eq!(expected, actual);
967 }
968
969 #[test]
970 fn a_star_matches_anything_that_exists() {
971 let expected: bool = true;
972 let actual: bool = if_none_match(&headers("*"), ETAG);
973 assert_eq!(expected, actual);
974 }
975
976 #[test]
977 fn one_match_anywhere_in_the_list_is_enough() {
978 let expected: bool = true;
982 let actual: bool = if_none_match(&headers("\"other\", \"abc123\""), ETAG);
983 assert_eq!(expected, actual);
984 }
985
986 #[test]
987 fn a_stale_validator_gets_the_new_document() {
988 let expected: bool = false;
989 let actual: bool = if_none_match(&headers("\"stale\""), ETAG);
990 assert_eq!(expected, actual);
991 }
992 }
993 use std::sync::Arc;
994 use std::sync::atomic::{AtomicBool, Ordering};
995 use std::time::Duration;
996
997 use axum::Router;
998 use axum::routing::get;
999 use tokio::time::timeout;
1000
1001 use super::super::shutdown::Shutdown;
1002 use super::{
1003 API_PREFIX, AppShutdown, DEFAULT_BODY_LIMIT, DEFAULT_DRAIN_TIMEOUT, WebServerError, health,
1004 serve_on,
1005 };
1006
1007 #[derive(Clone)]
1010 struct RecordingState {
1011 ran: Arc<AtomicBool>,
1012 linger: Option<Duration>,
1013 }
1014
1015 impl RecordingState {
1016 fn instant() -> Self {
1017 Self {
1018 ran: Arc::new(AtomicBool::new(false)),
1019 linger: None,
1020 }
1021 }
1022 }
1023
1024 impl AppShutdown for RecordingState {
1025 async fn on_shutdown(&self) {
1026 if let Some(linger) = self.linger {
1027 tokio::time::sleep(linger).await;
1028 }
1029 self.ran.store(true, Ordering::SeqCst);
1030 }
1031 }
1032
1033 async fn serve_until_shutdown(
1035 state: RecordingState,
1036 shutdown: Shutdown,
1037 ) -> Result<(), WebServerError> {
1038 let router: Router = Router::new().route("/health", get(health));
1039 let cleanup: RecordingState = state.clone();
1040 serve_on(router, "127.0.0.1", 0, shutdown, async move {
1042 cleanup.on_shutdown().await;
1043 })
1044 .await
1045 }
1046
1047 #[tokio::test]
1048 async fn app_cleanup_runs_when_the_shutdown_signal_arrives() {
1049 let shutdown: Shutdown = Shutdown::manual();
1050 let state: RecordingState = RecordingState::instant();
1051 let ran: Arc<AtomicBool> = Arc::clone(&state.ran);
1052
1053 shutdown.trigger();
1054 timeout(
1055 Duration::from_secs(1),
1056 serve_until_shutdown(state, shutdown.clone()),
1057 )
1058 .await
1059 .expect("serving ended promptly")
1060 .expect("serving ended cleanly");
1061
1062 let expected: bool = true;
1063 let actual: bool = ran.load(Ordering::SeqCst);
1064 assert_eq!(expected, actual);
1065 }
1066
1067 #[tokio::test]
1068 async fn app_cleanup_never_runs_when_the_server_dies_before_any_signal() {
1069 let shutdown: Shutdown = Shutdown::manual();
1073 let state: RecordingState = RecordingState::instant();
1074 let ran: Arc<AtomicBool> = Arc::clone(&state.ran);
1075 let cleanup: RecordingState = state.clone();
1076
1077 let router: Router = Router::new().route("/health", get(health));
1078 let result: Result<(), WebServerError> = timeout(
1079 Duration::from_secs(1),
1080 serve_on(router, "not-an-ip", 0, shutdown, async move {
1081 cleanup.on_shutdown().await;
1082 }),
1083 )
1084 .await
1085 .expect("a bind failure returns immediately, it does not wait for cleanup");
1086
1087 assert!(result.is_err(), "an unparseable host is a bind error");
1088
1089 let expected: bool = false;
1090 let actual: bool = ran.load(Ordering::SeqCst);
1091 assert_eq!(expected, actual);
1092 }
1093
1094 #[tokio::test(start_paused = true)]
1095 async fn a_cleanup_that_overruns_the_window_is_cancelled_rather_than_awaited() {
1096 let shutdown: Shutdown = Shutdown::manual();
1100 let state: RecordingState = RecordingState {
1101 ran: Arc::new(AtomicBool::new(false)),
1102 linger: Some(DEFAULT_DRAIN_TIMEOUT * 2),
1103 };
1104 let ran: Arc<AtomicBool> = Arc::clone(&state.ran);
1105
1106 shutdown.trigger();
1107 serve_until_shutdown(state, shutdown.clone())
1108 .await
1109 .expect("serving still ends cleanly when cleanup is cancelled");
1110
1111 let expected: bool = false;
1112 let actual: bool = ran.load(Ordering::SeqCst);
1113 assert_eq!(
1114 expected, actual,
1115 "the hook was cancelled at its await, so it never reached its final store"
1116 );
1117 }
1118
1119 #[tokio::test]
1120 async fn a_server_with_nothing_in_flight_stops_at_once_instead_of_waiting_out_the_window() {
1121 let shutdown: Shutdown = Shutdown::manual();
1122 let router: Router = Router::new().route("/health", get(health));
1123
1124 let serving: tokio::task::JoinHandle<Result<(), WebServerError>> = tokio::spawn({
1126 let shutdown: Shutdown = shutdown.clone();
1127 async move { serve_on(router, "127.0.0.1", 0, shutdown, async {}).await }
1128 });
1129
1130 shutdown.trigger();
1131
1132 timeout(Duration::from_secs(1), serving)
1135 .await
1136 .expect("the server returned as soon as the drain began")
1137 .expect("the serving task did not panic")
1138 .expect("serving ended cleanly");
1139 }
1140
1141 #[test]
1142 fn the_body_limit_accommodates_an_ordinary_form_post() {
1143 let expected: usize = 256 * 1024;
1146 let actual: usize = DEFAULT_BODY_LIMIT;
1147 assert_eq!(expected, actual);
1148 }
1149
1150 #[cfg(feature = "pages")]
1151 #[test]
1152 fn stripping_the_prefix_leaves_a_nestable_path() {
1153 let expected: String = String::from("/boggledygook-a3f2c1d8");
1154 let actual: String = super::strip_api_prefix("/api/v1/boggledygook-a3f2c1d8");
1155 assert_eq!(expected, actual);
1156 }
1157
1158 #[test]
1159 fn the_api_prefix_is_the_one_every_project_shares() {
1160 assert_eq!("/api/v1", API_PREFIX);
1161 }
1162}