diff --git a/Cargo.lock b/Cargo.lock index cbcd9ff..a05a489 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2812,6 +2812,7 @@ dependencies = [ "tracing-subscriber", "url", "urlencoding", + "valuable", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 1cda04e..4eb9611 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -11,8 +11,8 @@ sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "sqlite", "migra dotenvy = "0.15.7" serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" -tracing = "0.1" -tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt", "json"] } +tracing = { version = "0.1.44" } +tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt", "json", "serde", ] } hyper-util = { version = "0.1.19", features = ["client", "client-legacy"] } tower-http = { version = "0.6", features = ["cors", "compression-zstd", "trace"] } tower_governor = { version = "0.8.0", features = ["axum", "tracing"] } @@ -40,3 +40,4 @@ html-escape = "0.2.13" josekit = "0.10.3" dashmap = "6.1" tower = "0.5" +valuable = "0.1.1" diff --git a/src/main.rs b/src/main.rs index 1dd0fb5..fa7eef1 100644 --- a/src/main.rs +++ b/src/main.rs @@ -31,12 +31,14 @@ use std::{env, net::SocketAddr}; use tower_governor::{ GovernorLayer, governor::GovernorConfigBuilder, key_extractor::SmartIpKeyExtractor, }; +use tower_http::cors::AllowHeaders; +use tower_http::trace::{DefaultOnRequest, HttpMakeClassifier}; use tower_http::{ compression::CompressionLayer, cors::{Any, CorsLayer}, trace::TraceLayer, }; -use tracing::log; +use tracing::{Span, log}; use tracing_subscriber::{EnvFilter, fmt, prelude::*}; mod auth; @@ -200,7 +202,6 @@ ______________| | || / \ / \||/ \ / \ || | |______________ #[tokio::main] async fn main() -> Result<(), Box> { - setup_tracing(); let pds_env_location = env::var("PDS_ENV_LOCATION").unwrap_or_else(|_| "/pds/pds.env".to_string()); @@ -210,6 +211,8 @@ async fn main() -> Result<(), Box> { "Error loading pds.env file (ignore if you loaded your variables in the environment somehow else): {e}" ); } + // Sets up after the pds.env file is loaded + setup_tracing(); let pds_root = env::var("PDS_DATA_DIRECTORY").expect("PDS_DATA_DIRECTORY is not set in your pds.env file"); @@ -390,40 +393,14 @@ async fn main() -> Result<(), Box> { .map(|v| v.eq_ignore_ascii_case("true") || v == "1") .unwrap_or(false); - let app = if request_logging { - app.layer(TraceLayer::new_for_http() - .make_span_with(|req: &axum::http::Request| { - let headers: std::collections::HashMap<&str, Vec<&str>> = req.headers() - .keys() - .map(|k| { - let vals: Vec<&str> = req.headers() - .get_all(k) - .iter() - .filter_map(|v| v.to_str().ok()) - .collect(); - (k.as_str(), vals) - }) - .collect(); - let headers_json = serde_json::to_string(&headers).unwrap_or_default(); - - tracing::info_span!("request", - method = %req.method(), - path = %req.uri().path(), - headers = %headers_json, - ) - }) - .on_response(|resp: &axum::http::Response, latency: Duration, _span: &tracing::Span| { - tracing::info!(status = resp.status().as_u16(), latency_ms = latency.as_millis() as u64, "response"); - }) - ) + if request_logging { + app = app.layer(request_trace_layer()); + } + + let app = app .layer(CompressionLayer::new()) .layer(cors) - .with_state(state) - } else { - app.layer(CompressionLayer::new()) - .layer(cors) - .with_state(state) - }; + .with_state(state); let host = env::var("GATEKEEPER_HOST").unwrap_or_else(|_| "0.0.0.0".to_string()); let port: u16 = env::var("GATEKEEPER_PORT") @@ -493,3 +470,29 @@ async fn shutdown_signal() { _ = terminate => {}, } } + +fn request_trace_layer() -> TraceLayer< + HttpMakeClassifier, + impl Fn(&axum::http::Request) -> Span + Clone, + DefaultOnRequest, + impl Fn(&axum::http::Response, Duration, &Span) + Clone, +> { + TraceLayer::new_for_http() + .make_span_with(|req: &axum::http::Request| { + let headers = req.headers(); + tracing::info_span!("request", + method = %req.method(), + path = %req.uri().path(), + headers = %format!("{:?}", headers), + ) + }) + .on_response( + |resp: &axum::http::Response, latency: Duration, _span: &tracing::Span| { + tracing::info!( + status = resp.status().as_u16(), + latency_ms = latency.as_millis() as u64, + "response" + ); + }, + ) +}