use axum::{
Router,
extract::{Request, State},
http::{HeaderMap, StatusCode},
middleware::{self, Next},
response::{IntoResponse, Response},
};
use revoke_trace::{
propagator::{HttpHeaders, TracePropagator},
span::{SpanBuilder, SpanExt, SpanKind, SpanStatus},
};
use std::sync::Arc;
use tracing::{error, info, warn};
#[derive(Clone)]
pub struct TracingConfig {
pub service_name: String,
pub log_request_body: bool,
pub log_response_body: bool,
pub ignored_paths: Vec<String>,
}
impl Default for TracingConfig {
fn default() -> Self {
Self {
service_name: "axum-service".to_string(),
log_request_body: false,
log_response_body: false,
ignored_paths: vec!["/health".to_string(), "/metrics".to_string()],
}
}
}
#[derive(Clone)]
pub struct TracingState {
config: Arc<TracingConfig>,
propagator: Arc<TracePropagator>,
}
impl TracingState {
pub fn new(config: TracingConfig) -> Self {
Self {
config: Arc::new(config),
propagator: Arc::new(TracePropagator::w3c()),
}
}
}
pub fn tracing_layer(
config: TracingConfig,
) -> axum::middleware::FromFnLayer<
impl Fn(
Request,
Next,
)
-> std::pin::Pin<Box<dyn std::future::Future<Output = Result<Response, StatusCode>> + Send>>
+ Clone,
> {
let state = TracingState::new(config);
middleware::from_fn(move |req: Request, next: Next| {
let state = state.clone();
Box::pin(trace_request(state, req, next))
})
}
async fn trace_request(
state: TracingState,
req: Request,
next: Next,
) -> Result<Response, StatusCode> {
let uri = req.uri().clone();
let method = req.method().clone();
let path = uri.path();
if state.config.ignored_paths.iter().any(|p| path == p) {
return Ok(next.run(req).await);
}
let headers = extract_headers(req.headers());
let parent_context = state.propagator.extract(&headers);
let span_name = format!("{} {}", method, path);
let mut span_builder = SpanBuilder::new(&span_name)
.with_kind(SpanKind::Server)
.with_attribute("service.name", &state.config.service_name)
.with_attribute("http.method", method.as_str())
.with_attribute("http.url", &uri.to_string())
.with_attribute("http.target", path)
.with_attribute("http.scheme", uri.scheme_str().unwrap_or("http"))
.with_attribute("user_agent", extract_user_agent(req.headers()));
if let Some(host) = req.headers().get("host").and_then(|h| h.to_str().ok()) {
span_builder = span_builder.with_attribute("http.host", host);
}
if let Some(parent) = &parent_context {
span_builder = span_builder
.with_attribute("parent.trace_id", &parent.trace_id().to_string())
.with_attribute("parent.span_id", &parent.span_id().to_string());
}
let span = span_builder.start();
let _guard = span.enter();
info!(
method = %method,
path = %path,
"Request started"
);
let start_time = std::time::Instant::now();
let response = next.run(req).await;
let duration = start_time.elapsed();
let status = response.status();
span.record("http.status_code", &status.as_u16().to_string());
span.record("http.response_time_ms", &duration.as_millis().to_string());
if status.is_success() {
span.set_status(SpanStatus::ok());
info!(
status = %status,
duration_ms = %duration.as_millis(),
"Request completed successfully"
);
} else if status.is_client_error() {
span.set_status(SpanStatus::error(format!("Client error: {}", status)));
warn!(
status = %status,
duration_ms = %duration.as_millis(),
"Client error"
);
} else if status.is_server_error() {
span.set_status(SpanStatus::error(format!("Server error: {}", status)));
error!(
status = %status,
duration_ms = %duration.as_millis(),
"Server error"
);
}
Ok(response)
}
pub async fn error_handler(req: Request, next: Next) -> Result<Response, StatusCode> {
let response = next.run(req).await;
if response.status().is_server_error() {
tracing::Span::current().record("error", "true");
tracing::Span::current().record("error.type", "internal_server_error");
}
Ok(response)
}
fn extract_headers(headers: &HeaderMap) -> HttpHeaders {
headers
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|v| (name.to_string(), v.to_string()))
})
.collect()
}
fn extract_user_agent(headers: &HeaderMap) -> &str {
headers
.get("user-agent")
.and_then(|h| h.to_str().ok())
.unwrap_or("unknown")
}
pub fn inject_trace_headers(mut response: Response) -> Response {
let current_context = tracing::Span::current();
response
}
pub fn create_traced_router(config: TracingConfig) -> Router {
Router::new()
.layer(tracing_layer(config))
.layer(middleware::from_fn(error_handler))
}
#[cfg(test)]
mod tests {
use super::*;
use axum::routing::get;
#[tokio::test]
async fn test_tracing_middleware() {
let app = Router::new()
.route("/", get(|| async { "Hello, World!" }))
.layer(tracing_layer(TracingConfig::default()));
}
}