revoke-trace 0.3.0

Distributed tracing with OpenTelemetry for Revoke framework
Documentation
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);

    // 创建 span
    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());

    // 设置 span 状态
    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() {
        // 获取当前 span 并记录错误
        tracing::Span::current().record("error", "true");
        tracing::Span::current().record("error.type", "internal_server_error");
    }

    Ok(response)
}

/// 从 HeaderMap 提取 HttpHeaders
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()
}

/// 提取 User-Agent
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();

    // 这里可以添加自定义的追踪头
    // 例如:X-Trace-Id, X-Request-Id 等

    response
}

/// 创建带追踪的 Axum 路由器
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()));

        // 测试逻辑...
    }
}