use axum::extract::Request;
use axum::http::{HeaderMap, HeaderValue};
use axum::middleware::Next;
use axum::response::Response;
use std::collections::HashMap;
use std::sync::Arc;
use sz_orm_tracing::{Span, Tracer};
#[derive(Clone)]
pub struct TraceConfig {
pub tracer: Arc<dyn Tracer + Send + Sync>,
pub service_name: String,
pub exclude_paths: Vec<String>,
}
impl std::fmt::Debug for TraceConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TraceConfig")
.field("service_name", &self.service_name)
.field("exclude_paths", &self.exclude_paths)
.finish_non_exhaustive()
}
}
impl TraceConfig {
pub fn new(tracer: Arc<dyn Tracer + Send + Sync>) -> Self {
let service_name = "sz-rust".to_string();
Self {
tracer,
service_name,
exclude_paths: Vec::new(),
}
}
pub fn with_service_name(mut self, name: impl Into<String>) -> Self {
self.service_name = name.into();
self
}
pub fn with_exclude_paths(mut self, paths: Vec<String>) -> Self {
self.exclude_paths = paths;
self
}
pub fn is_excluded(&self, path: &str) -> bool {
crate::middleware::auth::is_route_allowed(path, &self.exclude_paths)
}
}
fn headers_to_hashmap(headers: &HeaderMap) -> HashMap<String, String> {
let mut map = HashMap::new();
for (name, value) in headers.iter() {
if let Ok(v) = value.to_str() {
map.insert(name.as_str().to_lowercase(), v.to_string());
}
}
map
}
pub fn extract_or_create_span(headers: &HeaderMap, config: &TraceConfig) -> Span {
let headers_map = headers_to_hashmap(headers);
let mut span = config
.tracer
.start_span(&format!("{}:request", config.service_name));
span.service_name = config.service_name.clone();
if let Some(parent_span) = config.tracer.extract(&headers_map) {
span.trace_id = parent_span.trace_id.clone();
span.parent_id = Some(parent_span.span_id.clone());
}
span
}
pub fn inject_traceparent_to_response(response: &mut Response, span: &Span, config: &TraceConfig) {
let headers_map = config.tracer.inject(span);
let headers = response.headers_mut();
for (key, value) in headers_map {
if let (Ok(name), Ok(header_value)) = (
axum::http::HeaderName::from_bytes(key.as_bytes()),
HeaderValue::from_str(&value),
) {
headers.insert(name, header_value);
}
}
}
pub async fn trace_middleware(
axum::extract::State(config): axum::extract::State<TraceConfig>,
req: Request,
next: Next,
) -> Response {
let path = req.uri().path().to_string();
if config.is_excluded(&path) {
return next.run(req).await;
}
let method = req.method().clone();
let uri = req.uri().path().to_string();
let mut span = extract_or_create_span(req.headers(), &config)
.with_tag("http.method", method.as_str())
.with_tag("http.uri", &uri)
.with_tag("http.path", &path);
let mut req = req;
req.extensions_mut().insert(span.clone());
let mut response = next.run(req).await;
let status = response.status().as_u16();
span = span.with_tag("http.status_code", status.to_string());
config.tracer.end_span(span.clone());
inject_traceparent_to_response(&mut response, &span, &config);
response
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::Router;
use http_body_util::BodyExt;
use tower::ServiceExt;
async fn read_body(resp: Response) -> String {
let bytes = resp.into_body().collect().await.unwrap().to_bytes();
String::from_utf8(bytes.to_vec()).unwrap()
}
fn make_request(method: &str, uri: &str) -> Request {
Request::builder()
.method(method)
.uri(uri)
.body(Body::empty())
.unwrap()
}
fn make_request_with_traceparent(method: &str, uri: &str, traceparent: &str) -> Request {
Request::builder()
.method(method)
.uri(uri)
.header("traceparent", traceparent)
.body(Body::empty())
.unwrap()
}
fn build_app() -> Router {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer).with_service_name("test-service");
Router::new()
.route(
"/api",
axum::routing::get(|| async { axum::http::StatusCode::OK }),
)
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
))
}
#[test]
fn test_trace_config_new() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer);
assert_eq!(config.service_name, "sz-rust");
assert!(config.exclude_paths.is_empty());
}
#[test]
fn test_trace_config_with_service_name() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("my-service");
assert_eq!(config.service_name, "my-service");
}
#[test]
fn test_trace_config_with_exclude_paths() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
assert_eq!(config.exclude_paths, vec!["/health".to_string()]);
}
#[test]
fn test_trace_config_is_excluded_exact_match() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
assert!(config.is_excluded("/health"));
assert!(!config.is_excluded("/api"));
}
#[test]
fn test_trace_config_is_excluded_wildcard_match() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/public/*".to_string()]);
assert!(config.is_excluded("/public/anything"));
assert!(!config.is_excluded("/api"));
}
#[test]
fn test_trace_config_is_excluded_empty_list() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer);
assert!(!config.is_excluded("/any"));
}
#[test]
fn test_trace_config_clone() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("cloned-service");
let cloned = config.clone();
assert_eq!(config.service_name, cloned.service_name);
}
#[test]
fn test_trace_config_debug() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("debug-service");
let debug_str = format!("{:?}", config);
assert!(debug_str.contains("debug-service"));
assert!(debug_str.contains("TraceConfig"));
}
#[test]
fn test_headers_to_hashmap_empty() {
let headers = HeaderMap::new();
let map = headers_to_hashmap(&headers);
assert!(map.is_empty());
}
#[test]
fn test_headers_to_hashmap_single_header() {
let mut headers = HeaderMap::new();
headers.insert("x-custom", "value1".parse().unwrap());
let map = headers_to_hashmap(&headers);
assert_eq!(map.get("x-custom"), Some(&"value1".to_string()));
}
#[test]
fn test_headers_to_hashmap_multiple_headers() {
let mut headers = HeaderMap::new();
headers.insert("x-custom-1", "value1".parse().unwrap());
headers.insert("x-custom-2", "value2".parse().unwrap());
let map = headers_to_hashmap(&headers);
assert_eq!(map.len(), 2);
assert_eq!(map.get("x-custom-1"), Some(&"value1".to_string()));
assert_eq!(map.get("x-custom-2"), Some(&"value2".to_string()));
}
#[test]
fn test_headers_to_hashmap_lowercases_keys() {
let mut headers = HeaderMap::new();
headers.insert("X-Custom", "value".parse().unwrap());
let map = headers_to_hashmap(&headers);
assert_eq!(map.get("x-custom"), Some(&"value".to_string()));
}
#[test]
fn test_headers_to_hashmap_skips_invalid_ascii() {
let mut headers = HeaderMap::new();
let invalid_value = HeaderValue::from_bytes(b"\xff\xfe").unwrap();
headers.insert("x-invalid", invalid_value);
let map = headers_to_hashmap(&headers);
assert!(!map.contains_key("x-invalid"));
}
#[test]
fn test_extract_or_create_span_no_traceparent_creates_new_span() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("my-service");
let headers = HeaderMap::new();
let span = extract_or_create_span(&headers, &config);
assert!(!span.trace_id().is_empty());
assert!(!span.span_id().is_empty());
assert!(span.parent_id().is_none());
assert_eq!(span.service_name(), "my-service");
}
#[test]
fn test_extract_or_create_span_with_traceparent_creates_child_span() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("my-service");
let parent_span = config.tracer.start_span("parent");
let parent_trace_id = parent_span.trace_id().to_string();
let parent_span_id = parent_span.span_id().to_string();
let headers_map = config.tracer.inject(&parent_span);
let mut headers = HeaderMap::new();
for (key, value) in &headers_map {
if let (Ok(name), Ok(header_value)) = (
axum::http::HeaderName::from_bytes(key.as_bytes()),
HeaderValue::from_str(value),
) {
headers.insert(name, header_value);
}
}
let child_span = extract_or_create_span(&headers, &config);
assert_eq!(child_span.trace_id(), parent_trace_id);
assert_eq!(child_span.parent_id(), Some(parent_span_id.as_str()));
}
#[test]
fn test_extract_or_create_span_with_invalid_traceparent_creates_new_span() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("my-service");
let mut headers = HeaderMap::new();
headers.insert("traceparent", "invalid".parse().unwrap());
let span = extract_or_create_span(&headers, &config);
assert!(span.parent_id().is_none());
}
#[test]
fn test_inject_traceparent_to_response_adds_headers() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("my-service");
let span = config.tracer.start_span("test");
let mut response = Response::new(Body::from("body"));
inject_traceparent_to_response(&mut response, &span, &config);
assert!(response.headers().contains_key("traceparent"));
}
#[test]
fn test_inject_traceparent_to_response_preserves_existing_headers() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("my-service");
let span = config.tracer.start_span("test");
let mut response = Response::builder()
.header("x-custom", "value")
.body(Body::from("body"))
.unwrap();
inject_traceparent_to_response(&mut response, &span, &config);
assert_eq!(
response
.headers()
.get("x-custom")
.unwrap()
.to_str()
.unwrap(),
"value"
);
assert!(response.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_creates_span_for_request() {
let app = build_app();
let resp = app.oneshot(make_request("GET", "/api")).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert!(resp.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_excluded_path_no_span() {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
let app = Router::new()
.route(
"/health",
axum::routing::get(|| async { axum::http::StatusCode::OK }),
)
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
));
let resp = app.oneshot(make_request("GET", "/health")).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert!(!resp.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_wildcard_exclude() {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/public/*".to_string()]);
let app = Router::new()
.route(
"/public/asset",
axum::routing::get(|| async { axum::http::StatusCode::OK }),
)
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
));
let resp = app
.oneshot(make_request("GET", "/public/asset"))
.await
.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert!(!resp.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_injects_span_into_extensions() {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer);
let app = Router::new()
.route(
"/api",
axum::routing::get(|| async { axum::http::StatusCode::OK }),
)
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
));
let resp = app.oneshot(make_request("GET", "/api")).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert!(resp.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_child_span_inherits_trace_id() {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer);
let parent_span = config.tracer.start_span("parent");
let parent_trace_id = parent_span.trace_id().to_string();
let headers_map = config.tracer.inject(&parent_span);
let traceparent = headers_map
.get("traceparent")
.expect("traceparent should be in injected headers");
let app = Router::new()
.route(
"/api",
axum::routing::get(|| async { axum::http::StatusCode::OK }),
)
.layer(axum::middleware::from_fn_with_state(
config.clone(),
trace_middleware,
));
let resp = app
.oneshot(make_request_with_traceparent("GET", "/api", traceparent))
.await
.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
let response_traceparent = resp.headers().get("traceparent").unwrap().to_str().unwrap();
let parts: Vec<&str> = response_traceparent.split('-').collect();
assert_eq!(parts.len(), 4);
assert_eq!(parts[1], parent_trace_id);
}
#[tokio::test]
async fn test_trace_middleware_preserves_response_body() {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer);
let app = Router::new()
.route("/body", axum::routing::get(|| async { "hello" }))
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
));
let resp = app.oneshot(make_request("GET", "/body")).await.unwrap();
let body = read_body(resp).await;
assert_eq!(body, "hello");
}
#[tokio::test]
async fn test_trace_middleware_handles_post_request() {
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer);
let app = Router::new()
.route(
"/submit",
axum::routing::post(|| async { axum::http::StatusCode::CREATED }),
)
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
));
let req = Request::builder()
.method("POST")
.uri("/submit")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::CREATED);
assert!(resp.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_chains_with_other_middleware() {
async fn add_header_middleware(req: Request, next: Next) -> Response {
let mut resp = next.run(req).await;
resp.headers_mut()
.insert("X-Custom", "value".parse().unwrap());
resp
}
let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
let config = TraceConfig::new(tracer);
let app = Router::new()
.route("/", axum::routing::get(|| async { "ok" }))
.layer(axum::middleware::from_fn(add_header_middleware))
.layer(axum::middleware::from_fn_with_state(
config,
trace_middleware,
));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), axum::http::StatusCode::OK);
assert_eq!(
resp.headers().get("X-Custom").unwrap().to_str().unwrap(),
"value"
);
assert!(resp.headers().contains_key("traceparent"));
}
#[tokio::test]
async fn test_trace_middleware_different_requests_different_trace_ids() {
let app = build_app();
let resp1 = app
.clone()
.oneshot(make_request("GET", "/api"))
.await
.unwrap();
let resp2 = app.oneshot(make_request("GET", "/api")).await.unwrap();
let tp1 = resp1
.headers()
.get("traceparent")
.unwrap()
.to_str()
.unwrap();
let tp2 = resp2
.headers()
.get("traceparent")
.unwrap()
.to_str()
.unwrap();
let parts1: Vec<&str> = tp1.split('-').collect();
let parts2: Vec<&str> = tp2.split('-').collect();
assert_ne!(parts1[1], parts2[1]); }
#[test]
fn test_php_session_init_no_tracing_capability() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer);
assert_eq!(config.service_name, "sz-rust");
}
#[test]
fn test_w3c_tracecontext_format_alignment() {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
let config = TraceConfig::new(tracer).with_service_name("test-service");
let span = config.tracer.start_span("test");
let headers_map = config.tracer.inject(&span);
let traceparent = headers_map
.get("traceparent")
.expect("traceparent should be present");
let parts: Vec<&str> = traceparent.split('-').collect();
assert_eq!(parts.len(), 4, "traceparent should have 4 parts");
assert_eq!(parts[0], "00", "version should be 00");
assert_eq!(parts[1].len(), 32, "trace_id should be 32 hex chars");
assert_eq!(parts[2].len(), 16, "span_id should be 16 hex chars");
assert_eq!(parts[3].len(), 2, "trace_flags should be 2 hex chars");
}
#[test]
fn test_trace_middleware_executes_first_in_order() {
use crate::middleware::order::{MiddlewareKind, DEFAULT_ORDER};
assert_eq!(DEFAULT_ORDER.first(), Some(&MiddlewareKind::Trace));
}
}