diff --git a/Cargo.lock b/Cargo.lock index 57ceed7fc..1948a481d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1823,6 +1823,7 @@ dependencies = [ "cc", "clap", "criterion", + "flate2", "fluxer_common", "futures-util", "hex", @@ -1850,6 +1851,7 @@ dependencies = [ "tokio", "tokio-util", "tower", + "tower-http 0.7.1", "tracing", "tracing-subscriber", "url", diff --git a/fluxer_media_proxy/Cargo.toml b/fluxer_media_proxy/Cargo.toml index 2408c4a27..5a8532b86 100644 --- a/fluxer_media_proxy/Cargo.toml +++ b/fluxer_media_proxy/Cargo.toml @@ -48,6 +48,7 @@ tempfile = "3.27.0" thiserror = "2.0.20" tokio = {version = "1.53.1", features = ["fs", "io-util", "macros", "net", "rt-multi-thread", "signal", "sync", "time"]} tokio-util = {version = "0.7.19", features = ["io"]} +tower-http = {version = "0.7.1", features = ["compression-gzip"]} tracing = "0.1.44" tracing-subscriber = {version = "0.3.23", features = ["fmt", "json"]} url = "2.5.8" @@ -59,5 +60,6 @@ pkg-config = "0.3.34" [dev-dependencies] criterion = "0.8.2" +flate2 = "1.1.10" proptest = "1.11.0" tower = {version = "0.5.3", features = ["util"]} diff --git a/fluxer_media_proxy/src/server/download_stream.rs b/fluxer_media_proxy/src/server/download_stream.rs index 25598ceef..210b8be3c 100644 --- a/fluxer_media_proxy/src/server/download_stream.rs +++ b/fluxer_media_proxy/src/server/download_stream.rs @@ -75,10 +75,13 @@ impl DownloadStreamPolicy { total_deadline: Instant::now().checked_add(total_timeout), overruns: self.overruns, }; - Body::from_stream(futures_util::stream::unfold( - Some(stream), - |stream: Option| async move { stream?.next().await }, - )) + Body::from_stream( + futures_util::stream::unfold( + Some(stream), + |stream: Option| async move { stream?.next().await }, + ) + .fuse(), + ) } } diff --git a/fluxer_media_proxy/src/server/runtime.rs b/fluxer_media_proxy/src/server/runtime.rs index 1ef6cc2e7..2e6cafd92 100644 --- a/fluxer_media_proxy/src/server/runtime.rs +++ b/fluxer_media_proxy/src/server/runtime.rs @@ -9,15 +9,18 @@ use super::{ }; use crate::{ aggregate_error::aggregate_results, - config::{Config, PolicyMode}, - media_process, request_log, + config::{Config, DeploymentMode, PolicyMode}, + media_process, mime, request_log, }; use axum::{ - Router, middleware, + Router, + http::{Extensions, HeaderMap, StatusCode, Version, header}, + middleware, routing::{any, get, post, put}, }; use std::{net::SocketAddr, sync::Arc, time::Duration}; use tokio::{net::TcpListener, time::timeout_at}; +use tower_http::compression::{CompressionLayer, CompressionLevel}; use tracing::info; pub async fn run(cfg: Config) -> anyhow::Result<()> { @@ -68,8 +71,44 @@ pub async fn run(cfg: Config) -> anyhow::Result<()> { drain_and_shutdown(&state, &drain).await } +const STATIC_COMPRESSIBLE_CONTENT_TYPES: [&str; 9] = [ + "application/javascript", + "text/javascript", + "text/css", + "text/html", + "application/json", + "application/manifest+json", + "application/xml", + "application/wasm", + "image/svg+xml", +]; + +const STATIC_COMPRESSION_MIN_BYTES: u64 = 1024; + +fn is_compressible_static_response( + status: StatusCode, + _version: Version, + headers: &HeaderMap, + _extensions: &Extensions, +) -> bool { + let content_type = headers + .get(header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()); + status == StatusCode::OK + && headers + .get(header::CONTENT_LENGTH) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse::().ok()) + .is_some_and(|length| length >= STATIC_COMPRESSION_MIN_BYTES) + && mime::normalize(content_type).is_some_and(|content_type| { + STATIC_COMPRESSIBLE_CONTENT_TYPES + .iter() + .any(|compressible| compressible.eq_ignore_ascii_case(content_type)) + }) +} + fn build_router(state: Arc, drain: HttpRequestDrain) -> Router { - Router::new() + let router = Router::new() .route("/_health", get(routes::ops::health)) .route("/_metrics", get(routes::ops::metrics_handler)) .route("/_metadata", post(routes::internal::metadata_handler)) @@ -80,7 +119,20 @@ fn build_router(state: Arc, drain: HttpRequestDrain) -> Router { "/v1/relay/{*key}", put(routes::relay::relay_put).options(routes::relay::relay_options), ) - .fallback(any(routes::dispatch::catch_all)) + .fallback(any(routes::dispatch::catch_all)); + let router = if state.cfg.mode == DeploymentMode::Static { + router.layer( + CompressionLayer::new() + .no_br() + .no_deflate() + .no_zstd() + .quality(CompressionLevel::Precise(5)) + .compress_when(is_compressible_static_response), + ) + } else { + router + }; + router .layer(middleware::from_fn(add_version_header)) .layer(middleware::from_fn_with_state( state.metrics.request(), @@ -142,7 +194,10 @@ async fn shutdown_signal() { #[cfg(test)] mod tests { use super::*; - use crate::http_headers; + use crate::{ + http_headers, + storage::tests::{FakeObject, FakeS3, fake_s3}, + }; use axum::{ body::Body, extract::ConnectInfo, @@ -325,6 +380,292 @@ mod tests { ); } + async fn static_router_serving(objects: &[(&str, &str, &[u8])]) -> (FakeS3, Router) { + let fake = fake_s3().await; + let tmp = tempfile::tempdir().expect("config root"); + let mut cfg = fake.config(tmp.path()); + cfg.mode = DeploymentMode::Static; + for (key, content_type, body) in objects { + fake.put_object( + &format!("{}/{key}", cfg.storage.bucket_static), + FakeObject { + body: body.to_vec(), + content_type: Some((*content_type).to_owned()), + ..FakeObject::default() + }, + ); + } + (fake, test_router_for(cfg)) + } + + async fn fetch( + router: Router, + path: &str, + accept_encoding: Option<&str>, + range: Option<&str>, + ) -> (Response, Vec) { + let mut builder = Request::builder().uri(path); + if let Some(value) = accept_encoding { + builder = builder.header(header::ACCEPT_ENCODING, value); + } + if let Some(value) = range { + builder = builder.header(header::RANGE, value); + } + let response = send(router, builder.body(Body::empty()).expect("request")).await; + let (parts, body) = response.into_parts(); + let bytes = axum::body::to_bytes(body, usize::MAX) + .await + .expect("response body") + .to_vec(); + (Response::from_parts(parts, Body::empty()), bytes) + } + + fn decoded(encoding: Option<&str>, body: &[u8]) -> Vec { + let mut out = Vec::new(); + match encoding { + Some("gzip") => { + std::io::Read::read_to_end(&mut flate2::read::GzDecoder::new(body), &mut out) + .expect("gzip body"); + } + None => out.extend_from_slice(body), + Some(other) => panic!("unexpected content-encoding {other}"), + } + out + } + + fn compressible_bytes(len: usize) -> Vec { + (0..len).map(|index| b"asm_module_"[index % 11]).collect() + } + + #[tokio::test] + async fn static_mode_gzips_scripts_styles_and_wasm_for_clients_that_accept_it() { + let body = compressible_bytes(64 * 1024); + let objects = [ + ("assets/0123456789abcdef.wasm", "application/wasm"), + ( + "assets/0123456789abcdef.js", + "application/javascript; charset=utf-8", + ), + ("assets/0123456789abcdef.css", "text/css; charset=utf-8"), + ( + "assets/0123456789abcdef.js.map", + "application/json; charset=utf-8", + ), + ("emoji/1f600.svg", "image/svg+xml"), + ]; + let stored: Vec<(&str, &str, &[u8])> = objects + .iter() + .map(|(key, content_type)| (*key, *content_type, &body[..])) + .collect(); + let (_fake, router) = static_router_serving(&stored).await; + for (key, _) in objects { + let path = format!("/{key}"); + for (accept, expected) in [ + ("br, gzip, zstd", Some("gzip")), + ("gzip, deflate", Some("gzip")), + ("gzip;q=0.5, br;q=1", Some("gzip")), + ("br", None), + ("deflate", None), + ("identity", None), + ("zstd", None), + ] { + let (response, bytes) = fetch(router.clone(), &path, Some(accept), None).await; + assert_eq!(StatusCode::OK, response.status(), "{key} {accept}"); + let encoding = header_text(&response, header::CONTENT_ENCODING); + assert_eq!(expected, encoding, "{key} {accept}"); + assert_eq!(body, decoded(encoding, &bytes), "{key} {accept}"); + assert_eq!( + vec!["Accept-Encoding"], + response + .headers() + .get_all(header::VARY) + .iter() + .map(|value| value.to_str().expect("vary is ascii")) + .collect::>(), + "{key} {accept}" + ); + if expected.is_some() { + assert!(bytes.len() < body.len() / 4, "{key} {accept}"); + assert_eq!(None, header_text(&response, header::CONTENT_LENGTH)); + assert_eq!(None, header_text(&response, header::ACCEPT_RANGES)); + } else { + assert_eq!( + Some(body.len().to_string().as_str()), + header_text(&response, header::CONTENT_LENGTH), + "{key} {accept}" + ); + } + } + let (response, bytes) = fetch(router.clone(), &path, None, None).await; + assert_eq!( + None, + header_text(&response, header::CONTENT_ENCODING), + "{key}" + ); + assert_eq!(body, bytes, "{key}"); + } + } + + #[tokio::test] + async fn static_mode_serves_archives_images_ranges_and_small_files_as_stored() { + let body = compressible_bytes(64 * 1024); + let small = compressible_bytes(512); + let (_fake, router) = static_router_serving(&[ + ("assets/0123456789abcdef.tar.gz", "application/gzip", &body), + ("avatars/0.png", "image/png", &body), + ( + "desktop/fluxer-setup.exe", + "application/octet-stream", + &body, + ), + ( + "assets/fedcba9876543210.js", + "application/javascript; charset=utf-8", + &small, + ), + ("assets/0123456789abcdef.wasm", "application/wasm", &body), + ]) + .await; + for (path, stored) in [ + ("/assets/0123456789abcdef.tar.gz", &body), + ("/avatars/0.png", &body), + ("/desktop/fluxer-setup.exe", &body), + ("/assets/fedcba9876543210.js", &small), + ] { + let (response, bytes) = fetch(router.clone(), path, Some("br, gzip"), None).await; + assert_eq!(StatusCode::OK, response.status(), "{path}"); + assert_eq!( + None, + header_text(&response, header::CONTENT_ENCODING), + "{path}" + ); + assert_eq!(stored, &bytes, "{path}"); + } + let (response, bytes) = fetch( + router, + "/assets/0123456789abcdef.wasm", + Some("br, gzip"), + Some("bytes=100-199"), + ) + .await; + assert_eq!(StatusCode::PARTIAL_CONTENT, response.status()); + assert_eq!(None, header_text(&response, header::CONTENT_ENCODING)); + assert_eq!(&body[100..200], &bytes[..]); + } + + #[tokio::test] + async fn static_head_matches_get_representation_headers_without_a_body() { + let body = compressible_bytes(1024); + let small = compressible_bytes(1023); + let (_fake, router) = static_router_serving(&[ + ("assets/0123456789abcdef.wasm", "application/wasm", &body), + ( + "assets/fedcba9876543210.js", + "application/javascript", + &small, + ), + ("assets/0123456789abcdef.tar.gz", "application/gzip", &body), + ]) + .await; + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("local listener"); + let address = listener.local_addr().expect("local address"); + let server = tokio::spawn(async move { + axum::serve( + listener, + router.into_make_service_with_connect_info::(), + ) + .await + .expect("local server"); + }); + let client = reqwest::Client::builder() + .no_proxy() + .no_gzip() + .no_brotli() + .no_deflate() + .no_zstd() + .build() + .expect("http client"); + for (path, accept_encoding, range, expected_encoding, expected_body) in [ + ( + "/assets/0123456789abcdef.wasm", + "gzip", + None, + Some("gzip"), + &body[..], + ), + ( + "/assets/0123456789abcdef.wasm", + "identity", + None, + None, + &body[..], + ), + ( + "/assets/fedcba9876543210.js", + "gzip", + None, + None, + &small[..], + ), + ( + "/assets/0123456789abcdef.tar.gz", + "gzip", + None, + None, + &body[..], + ), + ( + "/assets/0123456789abcdef.wasm", + "gzip", + Some("bytes=100-199"), + None, + &body[100..200], + ), + ] { + let url = format!("http://{address}{path}"); + let mut get = client + .get(&url) + .header(header::ACCEPT_ENCODING, accept_encoding); + let mut head = client + .head(&url) + .header(header::ACCEPT_ENCODING, accept_encoding); + if let Some(range) = range { + get = get.header(header::RANGE, range); + head = head.header(header::RANGE, range); + } + let get = get.send().await.expect("get response"); + let status = get.status(); + let get_headers = get.headers().clone(); + let get_body = get.bytes().await.expect("get body"); + let encoding = get_headers + .get(header::CONTENT_ENCODING) + .map(|value| value.to_str().expect("encoding is ascii")); + assert_eq!(expected_encoding, encoding, "{path} {accept_encoding}"); + assert_eq!(expected_body, decoded(encoding, &get_body)); + let head = head.send().await.expect("head response"); + assert_eq!(status, head.status(), "{path} {accept_encoding}"); + for name in [ + header::CONTENT_ENCODING, + header::CONTENT_LENGTH, + header::CONTENT_TYPE, + header::CONTENT_RANGE, + header::ACCEPT_RANGES, + header::VARY, + ] { + assert_eq!( + get_headers.get(&name), + head.headers().get(&name), + "{path} {accept_encoding} {name}" + ); + } + assert!(head.bytes().await.expect("head body").is_empty()); + } + server.abort(); + let _ = server.await; + } + #[tokio::test] async fn router_exposes_exactly_the_declared_route_and_method_table() { for mode in ["mp", "static", "upload", "relay"] {