1use axum::extract::Request;
55use axum::http::{HeaderMap, HeaderValue};
56use axum::middleware::Next;
57use axum::response::Response;
58use std::collections::HashMap;
59use std::sync::Arc;
60
61use crate::orm::{Span, Tracer};
62
63#[derive(Clone)]
65pub struct TraceConfig {
66 pub tracer: Arc<dyn Tracer + Send + Sync>,
68 pub service_name: String,
70 pub exclude_paths: Vec<String>,
72}
73
74impl std::fmt::Debug for TraceConfig {
75 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
76 f.debug_struct("TraceConfig")
77 .field("service_name", &self.service_name)
78 .field("exclude_paths", &self.exclude_paths)
79 .finish_non_exhaustive()
80 }
81}
82
83impl TraceConfig {
84 pub fn new(tracer: Arc<dyn Tracer + Send + Sync>) -> Self {
86 let service_name = "sz-rust".to_string();
87 Self {
88 tracer,
89 service_name,
90 exclude_paths: Vec::new(),
91 }
92 }
93
94 pub fn with_service_name(mut self, name: impl Into<String>) -> Self {
96 self.service_name = name.into();
97 self
98 }
99
100 pub fn with_exclude_paths(mut self, paths: Vec<String>) -> Self {
102 self.exclude_paths = paths;
103 self
104 }
105
106 pub fn is_excluded(&self, path: &str) -> bool {
108 crate::middleware::auth::is_route_allowed(path, &self.exclude_paths)
109 }
110}
111
112fn headers_to_hashmap(headers: &HeaderMap) -> HashMap<String, String> {
114 let mut map = HashMap::new();
115 for (name, value) in headers.iter() {
116 if let Ok(v) = value.to_str() {
117 map.insert(name.as_str().to_lowercase(), v.to_string());
118 }
119 }
120 map
121}
122
123pub fn extract_or_create_span(headers: &HeaderMap, config: &TraceConfig) -> Span {
134 let headers_map = headers_to_hashmap(headers);
135
136 let mut span = config
138 .tracer
139 .start_span(&format!("{}:request", config.service_name));
140 span.service_name = config.service_name.clone();
142
143 if let Some(parent_span) = config.tracer.extract(&headers_map) {
145 span.trace_id = parent_span.trace_id.clone();
146 span.parent_id = Some(parent_span.span_id.clone());
147 }
148 span
149}
150
151pub fn inject_traceparent_to_response(response: &mut Response, span: &Span, config: &TraceConfig) {
155 let headers_map = config.tracer.inject(span);
156 let headers = response.headers_mut();
157 for (key, value) in headers_map {
158 if let (Ok(name), Ok(header_value)) = (
160 axum::http::HeaderName::from_bytes(key.as_bytes()),
161 HeaderValue::from_str(&value),
162 ) {
163 headers.insert(name, header_value);
164 }
165 }
166}
167
168pub async fn trace_middleware(
185 axum::extract::State(config): axum::extract::State<TraceConfig>,
186 req: Request,
187 next: Next,
188) -> Response {
189 let path = req.uri().path().to_string();
190
191 if config.is_excluded(&path) {
193 return next.run(req).await;
194 }
195
196 let method = req.method().clone();
198 let uri = req.uri().path().to_string();
199 let mut span = extract_or_create_span(req.headers(), &config)
200 .with_tag("http.method", method.as_str())
201 .with_tag("http.uri", &uri)
202 .with_tag("http.path", &path);
203
204 let mut req = req;
206 req.extensions_mut().insert(span.clone());
207
208 let mut response = next.run(req).await;
210
211 let status = response.status().as_u16();
213 span = span.with_tag("http.status_code", status.to_string());
214 config.tracer.end_span(span.clone());
216
217 inject_traceparent_to_response(&mut response, &span, &config);
219
220 response
221}
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226 use axum::body::Body;
227 use axum::Router;
228 use http_body_util::BodyExt;
229 use tower::ServiceExt;
230
231 async fn read_body(resp: Response) -> String {
236 let bytes = resp.into_body().collect().await.unwrap().to_bytes();
237 String::from_utf8(bytes.to_vec()).unwrap()
238 }
239
240 fn make_request(method: &str, uri: &str) -> Request {
241 Request::builder()
242 .method(method)
243 .uri(uri)
244 .body(Body::empty())
245 .unwrap()
246 }
247
248 fn make_request_with_traceparent(method: &str, uri: &str, traceparent: &str) -> Request {
249 Request::builder()
250 .method(method)
251 .uri(uri)
252 .header("traceparent", traceparent)
253 .body(Body::empty())
254 .unwrap()
255 }
256
257 fn build_app() -> Router {
259 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
260 let config = TraceConfig::new(tracer).with_service_name("test-service");
261 Router::new()
262 .route(
263 "/api",
264 axum::routing::get(|| async { axum::http::StatusCode::OK }),
265 )
266 .layer(axum::middleware::from_fn_with_state(
267 config,
268 trace_middleware,
269 ))
270 }
271
272 #[test]
277 fn test_trace_config_new() {
278 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
279 let config = TraceConfig::new(tracer);
280 assert_eq!(config.service_name, "sz-rust");
281 assert!(config.exclude_paths.is_empty());
282 }
283
284 #[test]
285 fn test_trace_config_with_service_name() {
286 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
287 let config = TraceConfig::new(tracer).with_service_name("my-service");
288 assert_eq!(config.service_name, "my-service");
289 }
290
291 #[test]
292 fn test_trace_config_with_exclude_paths() {
293 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
294 let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
295 assert_eq!(config.exclude_paths, vec!["/health".to_string()]);
296 }
297
298 #[test]
299 fn test_trace_config_is_excluded_exact_match() {
300 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
301 let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
302 assert!(config.is_excluded("/health"));
303 assert!(!config.is_excluded("/api"));
304 }
305
306 #[test]
307 fn test_trace_config_is_excluded_wildcard_match() {
308 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
309 let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/public/*".to_string()]);
310 assert!(config.is_excluded("/public/anything"));
311 assert!(!config.is_excluded("/api"));
312 }
313
314 #[test]
315 fn test_trace_config_is_excluded_empty_list() {
316 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
317 let config = TraceConfig::new(tracer);
318 assert!(!config.is_excluded("/any"));
319 }
320
321 #[test]
322 fn test_trace_config_clone() {
323 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
324 let config = TraceConfig::new(tracer).with_service_name("cloned-service");
325 let cloned = config.clone();
326 assert_eq!(config.service_name, cloned.service_name);
327 }
328
329 #[test]
330 fn test_trace_config_debug() {
331 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
332 let config = TraceConfig::new(tracer).with_service_name("debug-service");
333 let debug_str = format!("{:?}", config);
334 assert!(debug_str.contains("debug-service"));
335 assert!(debug_str.contains("TraceConfig"));
336 }
337
338 #[test]
343 fn test_headers_to_hashmap_empty() {
344 let headers = HeaderMap::new();
345 let map = headers_to_hashmap(&headers);
346 assert!(map.is_empty());
347 }
348
349 #[test]
350 fn test_headers_to_hashmap_single_header() {
351 let mut headers = HeaderMap::new();
352 headers.insert("x-custom", "value1".parse().unwrap());
353 let map = headers_to_hashmap(&headers);
354 assert_eq!(map.get("x-custom"), Some(&"value1".to_string()));
355 }
356
357 #[test]
358 fn test_headers_to_hashmap_multiple_headers() {
359 let mut headers = HeaderMap::new();
360 headers.insert("x-custom-1", "value1".parse().unwrap());
361 headers.insert("x-custom-2", "value2".parse().unwrap());
362 let map = headers_to_hashmap(&headers);
363 assert_eq!(map.len(), 2);
364 assert_eq!(map.get("x-custom-1"), Some(&"value1".to_string()));
365 assert_eq!(map.get("x-custom-2"), Some(&"value2".to_string()));
366 }
367
368 #[test]
369 fn test_headers_to_hashmap_lowercases_keys() {
370 let mut headers = HeaderMap::new();
371 headers.insert("X-Custom", "value".parse().unwrap());
372 let map = headers_to_hashmap(&headers);
373 assert_eq!(map.get("x-custom"), Some(&"value".to_string()));
375 }
376
377 #[test]
378 fn test_headers_to_hashmap_skips_invalid_ascii() {
379 let mut headers = HeaderMap::new();
380 let invalid_value = HeaderValue::from_bytes(b"\xff\xfe").unwrap();
382 headers.insert("x-invalid", invalid_value);
383 let map = headers_to_hashmap(&headers);
384 assert!(!map.contains_key("x-invalid"));
386 }
387
388 #[test]
393 fn test_extract_or_create_span_no_traceparent_creates_new_span() {
394 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
395 let config = TraceConfig::new(tracer).with_service_name("my-service");
396 let headers = HeaderMap::new();
397 let span = extract_or_create_span(&headers, &config);
398 assert!(!span.trace_id().is_empty());
400 assert!(!span.span_id().is_empty());
401 assert!(span.parent_id().is_none());
403 assert_eq!(span.service_name(), "my-service");
405 }
406
407 #[test]
408 fn test_extract_or_create_span_with_traceparent_creates_child_span() {
409 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
410 let config = TraceConfig::new(tracer).with_service_name("my-service");
411
412 let parent_span = config.tracer.start_span("parent");
414 let parent_trace_id = parent_span.trace_id().to_string();
415 let parent_span_id = parent_span.span_id().to_string();
416 let headers_map = config.tracer.inject(&parent_span);
417
418 let mut headers = HeaderMap::new();
420 for (key, value) in &headers_map {
421 if let (Ok(name), Ok(header_value)) = (
422 axum::http::HeaderName::from_bytes(key.as_bytes()),
423 HeaderValue::from_str(value),
424 ) {
425 headers.insert(name, header_value);
426 }
427 }
428
429 let child_span = extract_or_create_span(&headers, &config);
430 assert_eq!(child_span.trace_id(), parent_trace_id);
432 assert_eq!(child_span.parent_id(), Some(parent_span_id.as_str()));
434 }
435
436 #[test]
437 fn test_extract_or_create_span_with_invalid_traceparent_creates_new_span() {
438 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
439 let config = TraceConfig::new(tracer).with_service_name("my-service");
440
441 let mut headers = HeaderMap::new();
442 headers.insert("traceparent", "invalid".parse().unwrap());
444
445 let span = extract_or_create_span(&headers, &config);
446 assert!(span.parent_id().is_none());
448 }
449
450 #[test]
455 fn test_inject_traceparent_to_response_adds_headers() {
456 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
457 let config = TraceConfig::new(tracer).with_service_name("my-service");
458 let span = config.tracer.start_span("test");
459
460 let mut response = Response::new(Body::from("body"));
461 inject_traceparent_to_response(&mut response, &span, &config);
462
463 assert!(response.headers().contains_key("traceparent"));
465 }
466
467 #[test]
468 fn test_inject_traceparent_to_response_preserves_existing_headers() {
469 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
470 let config = TraceConfig::new(tracer).with_service_name("my-service");
471 let span = config.tracer.start_span("test");
472
473 let mut response = Response::builder()
474 .header("x-custom", "value")
475 .body(Body::from("body"))
476 .unwrap();
477 inject_traceparent_to_response(&mut response, &span, &config);
478
479 assert_eq!(
481 response
482 .headers()
483 .get("x-custom")
484 .unwrap()
485 .to_str()
486 .unwrap(),
487 "value"
488 );
489 assert!(response.headers().contains_key("traceparent"));
491 }
492
493 #[tokio::test]
498 async fn test_trace_middleware_creates_span_for_request() {
499 let app = build_app();
500 let resp = app.oneshot(make_request("GET", "/api")).await.unwrap();
501 assert_eq!(resp.status(), axum::http::StatusCode::OK);
502 assert!(resp.headers().contains_key("traceparent"));
504 }
505
506 #[tokio::test]
507 async fn test_trace_middleware_excluded_path_no_span() {
508 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
509 let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/health".to_string()]);
510 let app = Router::new()
511 .route(
512 "/health",
513 axum::routing::get(|| async { axum::http::StatusCode::OK }),
514 )
515 .layer(axum::middleware::from_fn_with_state(
516 config,
517 trace_middleware,
518 ));
519
520 let resp = app.oneshot(make_request("GET", "/health")).await.unwrap();
521 assert_eq!(resp.status(), axum::http::StatusCode::OK);
522 assert!(!resp.headers().contains_key("traceparent"));
524 }
525
526 #[tokio::test]
527 async fn test_trace_middleware_wildcard_exclude() {
528 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
529 let config = TraceConfig::new(tracer).with_exclude_paths(vec!["/public/*".to_string()]);
530 let app = Router::new()
531 .route(
532 "/public/asset",
533 axum::routing::get(|| async { axum::http::StatusCode::OK }),
534 )
535 .layer(axum::middleware::from_fn_with_state(
536 config,
537 trace_middleware,
538 ));
539
540 let resp = app
541 .oneshot(make_request("GET", "/public/asset"))
542 .await
543 .unwrap();
544 assert_eq!(resp.status(), axum::http::StatusCode::OK);
545 assert!(!resp.headers().contains_key("traceparent"));
546 }
547
548 #[tokio::test]
549 async fn test_trace_middleware_injects_span_into_extensions() {
550 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
551 let config = TraceConfig::new(tracer);
552 let app = Router::new()
553 .route(
554 "/api",
555 axum::routing::get(|| async { axum::http::StatusCode::OK }),
556 )
557 .layer(axum::middleware::from_fn_with_state(
558 config,
559 trace_middleware,
560 ));
561
562 let resp = app.oneshot(make_request("GET", "/api")).await.unwrap();
564 assert_eq!(resp.status(), axum::http::StatusCode::OK);
565 assert!(resp.headers().contains_key("traceparent"));
566 }
567
568 #[tokio::test]
569 async fn test_trace_middleware_child_span_inherits_trace_id() {
570 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
571 let config = TraceConfig::new(tracer);
572
573 let parent_span = config.tracer.start_span("parent");
575 let parent_trace_id = parent_span.trace_id().to_string();
576 let headers_map = config.tracer.inject(&parent_span);
577 let traceparent = headers_map
578 .get("traceparent")
579 .expect("traceparent should be in injected headers");
580
581 let app = Router::new()
582 .route(
583 "/api",
584 axum::routing::get(|| async { axum::http::StatusCode::OK }),
585 )
586 .layer(axum::middleware::from_fn_with_state(
587 config.clone(),
588 trace_middleware,
589 ));
590
591 let resp = app
592 .oneshot(make_request_with_traceparent("GET", "/api", traceparent))
593 .await
594 .unwrap();
595 assert_eq!(resp.status(), axum::http::StatusCode::OK);
596
597 let response_traceparent = resp.headers().get("traceparent").unwrap().to_str().unwrap();
599 let parts: Vec<&str> = response_traceparent.split('-').collect();
601 assert_eq!(parts.len(), 4);
602 assert_eq!(parts[1], parent_trace_id);
604 }
605
606 #[tokio::test]
607 async fn test_trace_middleware_preserves_response_body() {
608 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
609 let config = TraceConfig::new(tracer);
610 let app = Router::new()
611 .route("/body", axum::routing::get(|| async { "hello" }))
612 .layer(axum::middleware::from_fn_with_state(
613 config,
614 trace_middleware,
615 ));
616
617 let resp = app.oneshot(make_request("GET", "/body")).await.unwrap();
618 let body = read_body(resp).await;
619 assert_eq!(body, "hello");
620 }
621
622 #[tokio::test]
623 async fn test_trace_middleware_handles_post_request() {
624 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
625 let config = TraceConfig::new(tracer);
626 let app = Router::new()
627 .route(
628 "/submit",
629 axum::routing::post(|| async { axum::http::StatusCode::CREATED }),
630 )
631 .layer(axum::middleware::from_fn_with_state(
632 config,
633 trace_middleware,
634 ));
635
636 let req = Request::builder()
637 .method("POST")
638 .uri("/submit")
639 .body(Body::empty())
640 .unwrap();
641 let resp = app.oneshot(req).await.unwrap();
642 assert_eq!(resp.status(), axum::http::StatusCode::CREATED);
643 assert!(resp.headers().contains_key("traceparent"));
644 }
645
646 #[tokio::test]
647 async fn test_trace_middleware_chains_with_other_middleware() {
648 async fn add_header_middleware(req: Request, next: Next) -> Response {
649 let mut resp = next.run(req).await;
650 resp.headers_mut()
651 .insert("X-Custom", "value".parse().unwrap());
652 resp
653 }
654
655 let tracer = Arc::new(sz_orm_tracing::SzTracer::new("test-service"));
656 let config = TraceConfig::new(tracer);
657 let app = Router::new()
658 .route("/", axum::routing::get(|| async { "ok" }))
659 .layer(axum::middleware::from_fn(add_header_middleware))
660 .layer(axum::middleware::from_fn_with_state(
661 config,
662 trace_middleware,
663 ));
664
665 let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
666 assert_eq!(resp.status(), axum::http::StatusCode::OK);
667 assert_eq!(
668 resp.headers().get("X-Custom").unwrap().to_str().unwrap(),
669 "value"
670 );
671 assert!(resp.headers().contains_key("traceparent"));
672 }
673
674 #[tokio::test]
675 async fn test_trace_middleware_different_requests_different_trace_ids() {
676 let app = build_app();
677 let resp1 = app
678 .clone()
679 .oneshot(make_request("GET", "/api"))
680 .await
681 .unwrap();
682 let resp2 = app.oneshot(make_request("GET", "/api")).await.unwrap();
683
684 let tp1 = resp1
685 .headers()
686 .get("traceparent")
687 .unwrap()
688 .to_str()
689 .unwrap();
690 let tp2 = resp2
691 .headers()
692 .get("traceparent")
693 .unwrap()
694 .to_str()
695 .unwrap();
696
697 let parts1: Vec<&str> = tp1.split('-').collect();
699 let parts2: Vec<&str> = tp2.split('-').collect();
700 assert_ne!(parts1[1], parts2[1]); }
702
703 #[test]
708 fn test_php_session_init_no_tracing_capability() {
709 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
713 let config = TraceConfig::new(tracer);
714 assert_eq!(config.service_name, "sz-rust");
716 }
717
718 #[test]
719 fn test_w3c_tracecontext_format_alignment() {
720 let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(sz_orm_tracing::SzTracer::new("test"));
723 let config = TraceConfig::new(tracer).with_service_name("test-service");
724 let span = config.tracer.start_span("test");
725 let headers_map = config.tracer.inject(&span);
726 let traceparent = headers_map
727 .get("traceparent")
728 .expect("traceparent should be present");
729
730 let parts: Vec<&str> = traceparent.split('-').collect();
732 assert_eq!(parts.len(), 4, "traceparent should have 4 parts");
733 assert_eq!(parts[0], "00", "version should be 00");
734 assert_eq!(parts[1].len(), 32, "trace_id should be 32 hex chars");
735 assert_eq!(parts[2].len(), 16, "span_id should be 16 hex chars");
736 assert_eq!(parts[3].len(), 2, "trace_flags should be 2 hex chars");
737 }
738
739 #[test]
740 fn test_trace_middleware_executes_first_in_order() {
741 use crate::middleware::order::{MiddlewareKind, DEFAULT_ORDER};
744 assert_eq!(DEFAULT_ORDER.first(), Some(&MiddlewareKind::Trace));
745 }
746}