Skip to main content

rest/
log.rs

1use actix_web::{
2    dev::{Service, ServiceRequest, ServiceResponse, Transform},
3    Error, HttpMessage,
4};
5use futures::future::{ok, LocalBoxFuture, Ready};
6use rust_zero_core::{LogContext, LogField, LogLevel, Logger, TraceContext};
7use std::{
8    task::{Context, Poll},
9    time::{Duration, Instant},
10};
11
12use crate::RequestIdValue;
13
14/// Emits a structured event for every completed HTTP request.
15pub struct LoggingMiddleware;
16
17impl<S, B> Transform<S, ServiceRequest> for LoggingMiddleware
18where
19    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
20    B: 'static,
21    S::Future: 'static,
22{
23    type Response = ServiceResponse<B>;
24    type Error = Error;
25    type Transform = LoggingMiddlewareService<S>;
26    type InitError = ();
27    type Future = Ready<Result<Self::Transform, Self::InitError>>;
28
29    fn new_transform(&self, service: S) -> Self::Future {
30        ok(LoggingMiddlewareService { service })
31    }
32}
33
34pub struct LoggingMiddlewareService<S> {
35    service: S,
36}
37
38impl<S, B> Service<ServiceRequest> for LoggingMiddlewareService<S>
39where
40    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
41    B: 'static,
42{
43    type Response = ServiceResponse<B>;
44    type Error = Error;
45    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
46
47    fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
48        self.service.poll_ready(cx)
49    }
50
51    fn call(&self, req: ServiceRequest) -> Self::Future {
52        let method = req.method().clone();
53        let path = req.path().to_owned();
54        let started_at = Instant::now();
55        let future = self.service.call(req);
56
57        Box::pin(async move {
58            match future.await {
59                Ok(response) => {
60                    tracing::info!(
61                        method = %method,
62                        path = %path,
63                        status = response.status().as_u16(),
64                        elapsed_ms = started_at.elapsed().as_millis() as u64,
65                        "HTTP request completed"
66                    );
67                    Ok(response)
68                }
69                Err(error) => {
70                    tracing::warn!(
71                        method = %method,
72                        path = %path,
73                        elapsed_ms = started_at.elapsed().as_millis() as u64,
74                        error = %error,
75                        "HTTP request failed"
76                    );
77                    Err(error)
78                }
79            }
80        })
81    }
82}
83
84/// Emits request logs through the standalone `rust-zero-core` structured logger.
85#[derive(Debug, Clone)]
86pub struct StructuredLogging {
87    logger: Logger,
88    slow_threshold: Option<Duration>,
89}
90
91impl StructuredLogging {
92    pub fn new(logger: Logger) -> Self {
93        Self {
94            logger,
95            slow_threshold: None,
96        }
97    }
98
99    /// Classifies completed calls at or above `threshold` as slow and adds stable transport-aware
100    /// fields suitable for log queries and alerts.
101    pub fn with_slow_threshold(mut self, threshold: Duration) -> Self {
102        assert!(!threshold.is_zero(), "slow-call threshold must be positive");
103        self.slow_threshold = Some(threshold);
104        self
105    }
106}
107
108impl<S, B> Transform<S, ServiceRequest> for StructuredLogging
109where
110    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
111    B: 'static,
112    S::Future: 'static,
113{
114    type Response = ServiceResponse<B>;
115    type Error = Error;
116    type Transform = StructuredLoggingService<S>;
117    type InitError = ();
118    type Future = Ready<Result<Self::Transform, Self::InitError>>;
119
120    fn new_transform(&self, service: S) -> Self::Future {
121        ok(StructuredLoggingService {
122            service,
123            logger: self.logger.clone(),
124            slow_threshold: self.slow_threshold,
125        })
126    }
127}
128
129pub struct StructuredLoggingService<S> {
130    service: S,
131    logger: Logger,
132    slow_threshold: Option<Duration>,
133}
134
135impl<S, B> Service<ServiceRequest> for StructuredLoggingService<S>
136where
137    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
138    B: 'static,
139{
140    type Response = ServiceResponse<B>;
141    type Error = Error;
142    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
143
144    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
145        self.service.poll_ready(context)
146    }
147
148    fn call(&self, request: ServiceRequest) -> Self::Future {
149        let method = request.method().to_string();
150        let path = request.path().to_owned();
151        let request_id = request
152            .extensions()
153            .get::<RequestIdValue>()
154            .map(|value| value.as_str().to_owned());
155        let trace = request.extensions().get::<TraceContext>().cloned();
156        let started_at = Instant::now();
157        let future = self.service.call(request);
158        let logger = self.logger.clone();
159        let slow_threshold = self.slow_threshold;
160
161        Box::pin(async move {
162            let context = request_context(request_id, trace);
163            match future.await {
164                Ok(response) => {
165                    let elapsed = started_at.elapsed();
166                    let slow = slow_threshold.is_some_and(|threshold| elapsed >= threshold);
167                    let mut fields = vec![
168                        LogField::new("transport", "http"),
169                        LogField::new("method", method),
170                        LogField::new("path", path),
171                        LogField::new("status", response.status().as_u16()),
172                        LogField::new("elapsed_ms", elapsed.as_millis() as u64),
173                        LogField::new("slow", slow),
174                    ];
175                    if let Some(route) = response.request().match_pattern() {
176                        fields.push(LogField::new("route", route));
177                    }
178                    if let Some(threshold) = slow_threshold {
179                        fields.push(LogField::new(
180                            "slow_threshold_ms",
181                            threshold.as_millis() as u64,
182                        ));
183                    }
184                    let _ = logger.log_with_context(
185                        if slow { LogLevel::Slow } else { LogLevel::Info },
186                        "HTTP request completed",
187                        Some(&context),
188                        fields,
189                    );
190                    Ok(response)
191                }
192                Err(error) => {
193                    let elapsed = started_at.elapsed();
194                    let slow = slow_threshold.is_some_and(|threshold| elapsed >= threshold);
195                    let mut fields = vec![
196                        LogField::new("transport", "http"),
197                        LogField::new("method", method),
198                        LogField::new("path", path),
199                        LogField::new("elapsed_ms", elapsed.as_millis() as u64),
200                        LogField::new("slow", slow),
201                        LogField::new("error", error.to_string()),
202                    ];
203                    if let Some(threshold) = slow_threshold {
204                        fields.push(LogField::new(
205                            "slow_threshold_ms",
206                            threshold.as_millis() as u64,
207                        ));
208                    }
209                    let _ = logger.log_with_context(
210                        LogLevel::Error,
211                        "HTTP request failed",
212                        Some(&context),
213                        fields,
214                    );
215                    Err(error)
216                }
217            }
218        })
219    }
220}
221
222fn request_context(request_id: Option<String>, trace: Option<TraceContext>) -> LogContext {
223    let mut context = LogContext::new();
224    if let Some(request_id) = request_id {
225        context = context.with_field(LogField::new("request_id", request_id));
226    }
227    if let Some(trace) = trace {
228        context = context.with_trace(trace);
229    }
230    context
231}
232
233#[cfg(test)]
234mod tests {
235    use super::*;
236    use crate::{RequestId, TraceContextMiddleware};
237    use actix_web::{test, web, App, HttpResponse};
238    use rust_zero_core::LogConfig;
239    use std::{
240        io::{self, Write},
241        sync::{Arc, Mutex},
242    };
243
244    #[derive(Clone, Default)]
245    struct SharedWriter(Arc<Mutex<Vec<u8>>>);
246
247    impl Write for SharedWriter {
248        fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
249            self.0.lock().unwrap().extend_from_slice(bytes);
250            Ok(bytes.len())
251        }
252
253        fn flush(&mut self) -> io::Result<()> {
254            Ok(())
255        }
256    }
257
258    #[actix_rt::test]
259    async fn emits_request_identity_and_trace_context() {
260        let output = SharedWriter::default();
261        let logger = Logger::to_writer(LogConfig::console("api"), output.clone()).unwrap();
262        let app = test::init_service(
263            App::new()
264                .wrap(StructuredLogging::new(logger))
265                .wrap(TraceContextMiddleware::new())
266                .wrap(RequestId::new())
267                .route(
268                    "/",
269                    web::get().to(|| async { HttpResponse::NoContent().finish() }),
270                ),
271        )
272        .await;
273
274        let response = test::call_service(
275            &app,
276            test::TestRequest::get()
277                .uri("/")
278                .insert_header(("x-request-id", "request-42"))
279                .insert_header((
280                    "traceparent",
281                    "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
282                ))
283                .to_request(),
284        )
285        .await;
286        assert_eq!(response.status(), 204);
287
288        let bytes = output.0.lock().unwrap().clone();
289        let record: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
290        assert_eq!(record["request_id"], "request-42");
291        assert_eq!(record["trace_id"], "4bf92f3577b34da6a3ce929d0e0e4736");
292        assert_eq!(record["method"], "GET");
293        assert_eq!(record["status"], 204);
294        assert_eq!(record["transport"], "http");
295        assert_eq!(record["route"], "/");
296        assert_eq!(record["slow"], false);
297    }
298
299    #[actix_rt::test]
300    async fn classifies_slow_requests_with_queryable_fields() {
301        let output = SharedWriter::default();
302        let logger = Logger::to_writer(LogConfig::console("api"), output.clone()).unwrap();
303        let app = test::init_service(
304            App::new()
305                .wrap(
306                    StructuredLogging::new(logger)
307                        .with_slow_threshold(std::time::Duration::from_millis(1)),
308                )
309                .route(
310                    "/users/{id}",
311                    web::get().to(|| async {
312                        actix_rt::time::sleep(std::time::Duration::from_millis(5)).await;
313                        HttpResponse::NoContent().finish()
314                    }),
315                ),
316        )
317        .await;
318
319        let response =
320            test::call_service(&app, test::TestRequest::get().uri("/users/42").to_request()).await;
321        assert_eq!(response.status(), 204);
322
323        let bytes = output.0.lock().unwrap().clone();
324        let record: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
325        assert_eq!(record["level"], "slow");
326        assert_eq!(record["route"], "/users/{id}");
327        assert_eq!(record["slow"], true);
328        assert_eq!(record["slow_threshold_ms"], 1);
329    }
330}