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
14pub 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#[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 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}