architect_sdk/middleware/
trace.rs1use axum::{
17 extract::Request,
18 http::{header::HeaderValue, HeaderName},
19 middleware::{from_fn, FromFnLayer, Next},
20 response::Response,
21};
22use std::future::Future;
23use std::pin::Pin;
24use tracing::Instrument;
25use uuid::Uuid;
26
27pub const TRACEPARENT_HEADER: &str = "traceparent";
29
30tokio::task_local! {
31 static CURRENT_TRACE_ID: String;
34}
35
36#[derive(Clone, Debug)]
39pub struct TraceId(pub String);
40
41type BoxFut = Pin<Box<dyn Future<Output = Response> + Send>>;
42type TraceFn = fn(Request, Next) -> BoxFut;
43
44pub fn trace_id_layer() -> FromFnLayer<TraceFn, (), (Request,)> {
52 from_fn(trace_id_mw as TraceFn)
53}
54
55pub fn current_trace_id() -> Option<String> {
60 CURRENT_TRACE_ID.try_with(|t| t.clone()).ok()
61}
62
63pub fn outbound_traceparent() -> Option<String> {
67 current_trace_id().map(|tid| format_traceparent(&tid))
68}
69
70fn format_traceparent(trace_id: &str) -> String {
73 let span_id = &Uuid::new_v4().simple().to_string()[..16];
74 format!("00-{trace_id}-{span_id}-01")
75}
76
77fn parse_trace_id(traceparent: &str) -> Option<String> {
81 let parts: [&str; 4] = {
82 let mut it = traceparent.trim().split('-');
83 let a = it.next()?;
84 let b = it.next()?;
85 let c = it.next()?;
86 let d = it.next()?;
87 if it.next().is_some() {
88 return None; }
90 [a, b, c, d]
91 };
92 let [version, trace_id, parent_id, flags] = parts;
93 let is_hex = |s: &str, n: usize| s.len() == n && s.bytes().all(|b| b.is_ascii_hexdigit());
94 if !is_hex(version, 2) || !is_hex(trace_id, 32) || !is_hex(parent_id, 16) || !is_hex(flags, 2) {
95 return None;
96 }
97 if version.eq_ignore_ascii_case("ff") {
98 return None;
99 }
100 let tid = trace_id.to_ascii_lowercase();
101 if tid.bytes().all(|b| b == b'0') {
102 return None;
103 }
104 Some(tid)
105}
106
107fn new_trace_id() -> String {
109 Uuid::new_v4().simple().to_string()
110}
111
112fn trace_id_mw(mut req: Request, next: Next) -> BoxFut {
113 let trace_id = req
116 .headers()
117 .get(TRACEPARENT_HEADER)
118 .and_then(|v| v.to_str().ok())
119 .and_then(parse_trace_id)
120 .unwrap_or_else(new_trace_id);
121
122 let method = req.method().clone();
123 let path = req.uri().path().to_string();
124
125 req.extensions_mut().insert(TraceId(trace_id.clone()));
127
128 let span = tracing::info_span!(
129 "request",
130 trace_id = %trace_id,
131 method = %method,
132 path = %path,
133 );
134 let response_tp = format_traceparent(&trace_id);
136
137 Box::pin(
138 CURRENT_TRACE_ID
139 .scope(trace_id, async move {
140 let mut response = next.run(req).await;
141 if let Ok(value) = HeaderValue::from_str(&response_tp) {
142 response
143 .headers_mut()
144 .insert(HeaderName::from_static(TRACEPARENT_HEADER), value);
145 }
146 response
147 })
148 .instrument(span),
149 )
150}
151
152#[cfg(test)]
153mod tests {
154 use super::*;
155 use axum::{body::Body, http::Request as HttpRequest, routing::get, Router};
156 use std::sync::{Arc, Mutex};
157 use tower::ServiceExt;
158
159 #[test]
160 fn parse_trace_id_accepts_valid_traceparent() {
161 let tp = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
162 assert_eq!(
163 parse_trace_id(tp).as_deref(),
164 Some("4bf92f3577b34da6a3ce929d0e0e4736")
165 );
166 }
167
168 #[test]
169 fn parse_trace_id_rejects_malformed_and_zero() {
170 assert_eq!(parse_trace_id("garbage"), None);
171 assert_eq!(parse_trace_id("00-abc-00f067aa0ba902b7-01"), None); assert_eq!(
173 parse_trace_id("ff-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"),
174 None
175 ); assert_eq!(
177 parse_trace_id("00-00000000000000000000000000000000-00f067aa0ba902b7-01"),
178 None
179 ); }
181
182 #[derive(Clone)]
185 struct BufWriter(Arc<Mutex<Vec<u8>>>);
186
187 impl std::io::Write for BufWriter {
188 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
189 self.0.lock().unwrap().extend_from_slice(buf);
190 Ok(buf.len())
191 }
192 fn flush(&mut self) -> std::io::Result<()> {
193 Ok(())
194 }
195 }
196
197 impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for BufWriter {
198 type Writer = BufWriter;
199 fn make_writer(&'a self) -> Self::Writer {
200 self.clone()
201 }
202 }
203
204 async fn handler() -> &'static str {
205 tracing::info!("handling request");
207 "ok"
208 }
209
210 fn app() -> Router {
211 Router::new()
212 .route("/x", get(handler))
213 .layer(trace_id_layer())
214 }
215
216 fn resp_trace_id(header: Option<&str>) -> Option<String> {
218 header.and_then(parse_trace_id)
219 }
220
221 #[test]
222 fn incoming_traceparent_flows_into_logs_and_response() {
223 let buf = Arc::new(Mutex::new(Vec::new()));
224 let subscriber = tracing_subscriber::fmt()
225 .json()
226 .flatten_event(true)
227 .with_current_span(true)
228 .with_span_list(true)
229 .with_max_level(tracing::Level::TRACE)
230 .with_writer(BufWriter(buf.clone()))
231 .finish();
232
233 tracing::subscriber::set_global_default(subscriber)
239 .expect("no other test should set a global subscriber");
240
241 let rt = tokio::runtime::Builder::new_current_thread()
242 .enable_all()
243 .build()
244 .unwrap();
245
246 let incoming = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
247 let echoed = rt.block_on(async {
248 let req = HttpRequest::builder()
249 .uri("/x")
250 .header(TRACEPARENT_HEADER, incoming)
251 .body(Body::empty())
252 .unwrap();
253 let resp = app().oneshot(req).await.unwrap();
254 resp.headers()
255 .get(TRACEPARENT_HEADER)
256 .and_then(|v| v.to_str().ok())
257 .map(String::from)
258 });
259
260 assert_eq!(
262 resp_trace_id(echoed.as_deref()).as_deref(),
263 Some("4bf92f3577b34da6a3ce929d0e0e4736")
264 );
265
266 let logs = String::from_utf8(buf.lock().unwrap().clone()).unwrap();
268 assert!(
269 logs.contains("\"trace_id\":\"4bf92f3577b34da6a3ce929d0e0e4736\""),
270 "expected trace_id in JSON logs, got: {logs}"
271 );
272 assert!(logs.contains("handling request"));
273 }
274
275 #[test]
276 fn missing_header_generates_root_trace() {
277 let rt = tokio::runtime::Builder::new_current_thread()
278 .enable_all()
279 .build()
280 .unwrap();
281 let echoed = rt.block_on(async {
282 let req = HttpRequest::builder()
283 .uri("/x")
284 .body(Body::empty())
285 .unwrap();
286 let resp = app().oneshot(req).await.unwrap();
287 resp.headers()
288 .get(TRACEPARENT_HEADER)
289 .and_then(|v| v.to_str().ok())
290 .map(String::from)
291 });
292 let tid = resp_trace_id(echoed.as_deref());
294 assert!(
295 tid.as_deref().map(|t| t.len() == 32).unwrap_or(false),
296 "expected a generated 32-hex trace id, got: {echoed:?}"
297 );
298 }
299}