kegani 0.1.0

A developer-friendly, ergonomic, production-ready Rust web framework
Documentation
//! Tracing middleware for OpenTelemetry integration
//!
//! Provides distributed tracing support for HTTP requests.

use actix_web::{
    dev::{Service, ServiceRequest, ServiceResponse, Transform},
    Error,
};
use std::future::{ready, Ready};
use std::pin::Pin;
use std::task::{Context, Poll};
use tracing::{span, Level};
use uuid::Uuid;

/// Tracing middleware configuration
#[derive(Debug, Clone)]
pub struct TracingConfig {
    pub service_name: String,
    pub trace_header: String,
}

impl Default for TracingConfig {
    fn default() -> Self {
        Self {
            service_name: "kegani".to_string(),
            trace_header: "x-trace-id".to_string(),
        }
    }
}

/// Tracing middleware
pub struct Tracing {
    config: TracingConfig,
}

impl Tracing {
    /// Create a new Tracing middleware
    pub fn new() -> Self {
        Self {
            config: TracingConfig::default(),
        }
    }

    /// Set the service name
    pub fn service_name(mut self, name: &str) -> Self {
        self.config.service_name = name.to_string();
        self
    }
}

impl Default for Tracing {
    fn default() -> Self {
        Self::new()
    }
}

impl<S, B> Transform<S, ServiceRequest> for Tracing
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = Error;
    type Transform = TracingMiddleware<S>;
    type InitError = ();
    type Future = Ready<Result<Self::Transform, Self::InitError>>;

    fn new_transform(&self, service: S) -> Self::Future {
        ready(Ok(TracingMiddleware {
            service,
            config: self.config.clone(),
        }))
    }
}

pub struct TracingMiddleware<S> {
    service: S,
    config: TracingConfig,
}

impl<S, B> Service<ServiceRequest> for TracingMiddleware<S>
where
    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
    S::Future: 'static,
    B: 'static,
{
    type Response = ServiceResponse<B>;
    type Error = Error;
    type Future = Pin<Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>>>>;

    fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
        self.service.poll_ready(cx)
    }

    fn call(&self, req: ServiceRequest) -> Self::Future {
        let trace_id = Uuid::new_v4().to_string();
        let trace_header = self.config.trace_header.clone();
        let service_name = self.config.service_name.clone();

        let span = span!(
            Level::INFO,
            "http_request",
            trace_id = %trace_id,
            method = %req.method(),
            path = %req.path(),
            service_name = %service_name
        );

        let _enter = span.enter();

        let fut = self.service.call(req);

        Box::pin(async move {
            let mut res = fut.await?;

            // Add trace ID to response headers
            let header_name = actix_web::http::header::HeaderName::try_from(trace_header.as_bytes())
                .unwrap_or_else(|_| actix_web::http::header::HeaderName::from_static("x-trace-id"));

            res.headers_mut().insert(
                header_name,
                trace_id.parse().unwrap(),
            );

            Ok(res.map_body(|_, body| body))
        })
    }
}