Skip to content

Commit 04e48d5

Browse files
authored
feat(server): add request-ID middleware for request correlation (#1082)
Add a UUID-based request-ID middleware using tower-http's request-id feature. Each inbound request receives a unique x-request-id header (or preserves a client-supplied one), which is recorded in the tracing span and propagated to the response. This enables operators to correlate log lines across the middleware stack for a single request under concurrent load, and lets clients reference specific requests in bug reports. Signed-off-by: sauagarwa <sauagarw@redhat.com>
1 parent 25c4fde commit 04e48d5

4 files changed

Lines changed: 338 additions & 45 deletions

File tree

‎Cargo.lock‎

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎Cargo.toml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ prost-types = "0.13"
2525
# HTTP server
2626
axum = { version = "0.8", features = ["ws"] }
2727
tower = "0.5"
28-
tower-http = { version = "0.6", features = ["cors", "trace"] }
28+
tower-http = { version = "0.6", features = ["cors", "trace", "request-id"] }
2929
hyper = { version = "1.6", features = ["full"] }
3030
hyper-util = { version = "0.1", features = ["tokio", "server-auto"] }
3131
http = "1.2"

‎crates/openshell-server/src/multiplex.rs‎

Lines changed: 263 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
//! to either the gRPC service or HTTP endpoints based on the request headers.
88
99
use bytes::Bytes;
10-
use http::{Request, Response};
10+
use http::{HeaderValue, Request, Response};
1111
use http_body::Body;
1212
use http_body_util::BodyExt;
1313
use hyper::body::Incoming;
@@ -25,15 +25,86 @@ use std::sync::Arc;
2525
use std::task::{Context, Poll};
2626
use std::time::{Duration, Instant};
2727
use tokio::io::{AsyncRead, AsyncWrite};
28-
use tower::{ServiceBuilder, ServiceExt};
29-
use tower_http::trace::TraceLayer;
28+
use tower::ServiceExt;
29+
use tower_http::request_id::{MakeRequestId, RequestId};
3030
use tracing::Span;
3131

3232
use crate::{
3333
OpenShellService, ServerState, auth::authz::AuthzPolicy, auth::oidc, http_router,
3434
inference::InferenceService,
3535
};
3636

37+
/// Request-ID generator that produces a UUID v4 for each inbound request.
38+
#[derive(Clone)]
39+
struct UuidRequestId;
40+
41+
impl MakeRequestId for UuidRequestId {
42+
fn make_request_id<B>(&mut self, _req: &Request<B>) -> Option<RequestId> {
43+
let id = uuid::Uuid::new_v4().to_string();
44+
Some(RequestId::new(HeaderValue::from_str(&id).unwrap()))
45+
}
46+
}
47+
48+
/// Build a tracing span for an inbound request, recording the `request_id`
49+
/// header (set by [`UuidRequestId`] or supplied by the client).
50+
fn make_request_span<B>(req: &Request<B>) -> Span {
51+
let path = req.uri().path();
52+
let request_id = req
53+
.headers()
54+
.get("x-request-id")
55+
.and_then(|v| v.to_str().ok())
56+
.unwrap_or("-");
57+
58+
if matches!(path, "/health" | "/healthz" | "/readyz") {
59+
tracing::debug_span!(
60+
"request",
61+
method = %req.method(),
62+
path,
63+
request_id,
64+
)
65+
} else {
66+
tracing::info_span!(
67+
"request",
68+
method = %req.method(),
69+
path,
70+
request_id,
71+
)
72+
}
73+
}
74+
75+
/// Log response status and latency within the request span.
76+
fn log_response<B>(res: &Response<B>, latency: Duration, _span: &Span) {
77+
tracing::info!(
78+
status = res.status().as_u16(),
79+
latency_ms = latency.as_millis(),
80+
"response"
81+
);
82+
}
83+
84+
/// Wrap a service with the standard request-ID middleware stack.
85+
///
86+
/// Layer order: `SetRequestId` → `TraceLayer` → `PropagateRequestId`.
87+
macro_rules! request_id_middleware {
88+
($service:expr) => {{
89+
let x_request_id = ::http::HeaderName::from_static("x-request-id");
90+
::tower::ServiceBuilder::new()
91+
.layer(::tower_http::request_id::SetRequestIdLayer::new(
92+
x_request_id.clone(),
93+
UuidRequestId,
94+
))
95+
.layer(
96+
::tower_http::trace::TraceLayer::new_for_http()
97+
.make_span_with(make_request_span)
98+
.on_request(())
99+
.on_response(log_response),
100+
)
101+
.layer(::tower_http::request_id::PropagateRequestIdLayer::new(
102+
x_request_id,
103+
))
104+
.service($service)
105+
}};
106+
}
107+
37108
/// Maximum inbound gRPC message size (1 MB).
38109
///
39110
/// Replaces tonic's implicit 4 MB default with a conservative limit to
@@ -77,22 +148,8 @@ impl MultiplexService {
77148
);
78149
let http_service = http_router(self.state.clone());
79150

80-
let grpc_service = ServiceBuilder::new()
81-
.layer(
82-
TraceLayer::new_for_http()
83-
.make_span_with(make_request_span)
84-
.on_request(())
85-
.on_response(log_response),
86-
)
87-
.service(grpc_service);
88-
let http_service = ServiceBuilder::new()
89-
.layer(
90-
TraceLayer::new_for_http()
91-
.make_span_with(make_request_span)
92-
.on_request(())
93-
.on_response(log_response),
94-
)
95-
.service(http_service);
151+
let grpc_service = request_id_middleware!(grpc_service);
152+
let http_service = request_id_middleware!(http_service);
96153

97154
let service = MultiplexedService::new(grpc_service, http_service);
98155

@@ -400,31 +457,6 @@ where
400457
}
401458
}
402459

403-
fn make_request_span<B>(req: &Request<B>) -> Span {
404-
let path = req.uri().path();
405-
if matches!(path, "/health" | "/healthz" | "/readyz") {
406-
tracing::debug_span!(
407-
"request",
408-
method = %req.method(),
409-
path,
410-
)
411-
} else {
412-
tracing::info_span!(
413-
"request",
414-
method = %req.method(),
415-
path,
416-
)
417-
}
418-
}
419-
420-
fn log_response<B>(res: &Response<B>, latency: Duration, _span: &Span) {
421-
tracing::info!(
422-
status = res.status().as_u16(),
423-
latency_ms = latency.as_millis(),
424-
"response"
425-
);
426-
}
427-
428460
fn grpc_method_from_path(path: &str) -> String {
429461
path.rsplit('/').next().unwrap_or(path).to_string()
430462
}
@@ -473,6 +505,193 @@ impl Body for BoxBody {
473505
#[cfg(test)]
474506
mod tests {
475507
use super::*;
508+
use bytes::Bytes;
509+
use http_body_util::Empty;
510+
use std::sync::Mutex;
511+
512+
#[test]
513+
fn uuid_request_id_generates_valid_uuid() {
514+
let mut maker = UuidRequestId;
515+
let req = Request::builder().body(()).unwrap();
516+
let id = maker.make_request_id(&req).expect("should produce an ID");
517+
let value = id.header_value().to_str().unwrap();
518+
uuid::Uuid::parse_str(value).expect("should be a valid UUID");
519+
}
520+
521+
#[test]
522+
fn uuid_request_id_generates_unique_ids() {
523+
let mut maker = UuidRequestId;
524+
let req = Request::builder().body(()).unwrap();
525+
let id1 = maker.make_request_id(&req).unwrap();
526+
let id2 = maker.make_request_id(&req).unwrap();
527+
assert_ne!(id1.header_value(), id2.header_value());
528+
}
529+
530+
async fn start_http_server_with_middleware() -> std::net::SocketAddr {
531+
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
532+
let addr = listener.local_addr().unwrap();
533+
534+
let http_service = crate::http::health_router();
535+
let http_service = request_id_middleware!(http_service);
536+
537+
let service = MultiplexedService::new(http_service.clone(), http_service);
538+
539+
tokio::spawn(async move {
540+
loop {
541+
let Ok((stream, _)) = listener.accept().await else {
542+
continue;
543+
};
544+
let svc = service.clone();
545+
tokio::spawn(async move {
546+
let _ = Builder::new(TokioExecutor::new())
547+
.serve_connection(TokioIo::new(stream), svc)
548+
.await;
549+
});
550+
}
551+
});
552+
553+
addr
554+
}
555+
556+
async fn http1_get(
557+
addr: std::net::SocketAddr,
558+
path: &str,
559+
headers: &[(&str, &str)],
560+
) -> Response<Incoming> {
561+
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
562+
let (mut sender, conn) = hyper::client::conn::http1::Builder::new()
563+
.handshake(TokioIo::new(stream))
564+
.await
565+
.unwrap();
566+
tokio::spawn(async move {
567+
let _ = conn.await;
568+
});
569+
570+
let mut builder = Request::builder()
571+
.method("GET")
572+
.uri(format!("http://{addr}{path}"));
573+
for (k, v) in headers {
574+
builder = builder.header(*k, *v);
575+
}
576+
let req = builder.body(Empty::<Bytes>::new()).unwrap();
577+
sender.send_request(req).await.unwrap()
578+
}
579+
580+
#[tokio::test]
581+
async fn http_response_includes_request_id() {
582+
let addr = start_http_server_with_middleware().await;
583+
let resp = http1_get(addr, "/healthz", &[]).await;
584+
assert_eq!(resp.status(), 200);
585+
586+
let request_id = resp
587+
.headers()
588+
.get("x-request-id")
589+
.expect("response should include x-request-id header");
590+
let id_str = request_id.to_str().unwrap();
591+
uuid::Uuid::parse_str(id_str).expect("should be a valid UUID");
592+
}
593+
594+
#[tokio::test]
595+
async fn http_preserves_client_request_id() {
596+
let addr = start_http_server_with_middleware().await;
597+
let client_id = "my-custom-correlation-id";
598+
let resp = http1_get(addr, "/healthz", &[("x-request-id", client_id)]).await;
599+
assert_eq!(resp.status(), 200);
600+
601+
let request_id = resp
602+
.headers()
603+
.get("x-request-id")
604+
.expect("response should include x-request-id header");
605+
assert_eq!(request_id.to_str().unwrap(), client_id);
606+
}
607+
608+
#[tokio::test]
609+
async fn each_request_gets_unique_id() {
610+
let addr = start_http_server_with_middleware().await;
611+
612+
let mut ids = Vec::new();
613+
for _ in 0..3 {
614+
let resp = http1_get(addr, "/healthz", &[]).await;
615+
let id = resp
616+
.headers()
617+
.get("x-request-id")
618+
.unwrap()
619+
.to_str()
620+
.unwrap()
621+
.to_string();
622+
ids.push(id);
623+
}
624+
625+
assert_ne!(ids[0], ids[1]);
626+
assert_ne!(ids[1], ids[2]);
627+
assert_ne!(ids[0], ids[2]);
628+
}
629+
630+
#[tokio::test]
631+
async fn grpc_path_includes_request_id() {
632+
let addr = start_http_server_with_middleware().await;
633+
let resp = http1_get(
634+
addr,
635+
"/openshell.v1.OpenShell/Health",
636+
&[
637+
("content-type", "application/grpc"),
638+
("x-request-id", "grpc-corr-id"),
639+
],
640+
)
641+
.await;
642+
643+
let request_id = resp
644+
.headers()
645+
.get("x-request-id")
646+
.expect("gRPC-routed response should include x-request-id header");
647+
assert_eq!(request_id.to_str().unwrap(), "grpc-corr-id");
648+
}
649+
650+
#[derive(Clone)]
651+
struct TraceBuf(Arc<Mutex<Vec<u8>>>);
652+
653+
impl std::io::Write for TraceBuf {
654+
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
655+
self.0.lock().unwrap().extend_from_slice(buf);
656+
Ok(buf.len())
657+
}
658+
659+
fn flush(&mut self) -> std::io::Result<()> {
660+
Ok(())
661+
}
662+
}
663+
664+
#[test]
665+
fn request_id_appears_in_trace_span() {
666+
use tracing_subscriber::fmt::format::FmtSpan;
667+
use tracing_subscriber::layer::SubscriberExt;
668+
669+
let log_buf: Arc<Mutex<Vec<u8>>> = Arc::new(Mutex::new(Vec::new()));
670+
let writer = TraceBuf(log_buf.clone());
671+
672+
let fmt_layer = tracing_subscriber::fmt::layer()
673+
.with_writer(move || writer.clone())
674+
.with_ansi(false)
675+
.with_span_events(FmtSpan::CLOSE);
676+
677+
let subscriber = tracing_subscriber::registry().with(fmt_layer);
678+
let _guard = tracing::subscriber::set_default(subscriber);
679+
680+
let req = Request::builder()
681+
.uri("/test-path")
682+
.header("x-request-id", "trace-test-id-12345")
683+
.body(Empty::<Bytes>::new())
684+
.unwrap();
685+
let span = make_request_span(&req);
686+
drop(span.enter());
687+
drop(span);
688+
689+
let output = String::from_utf8(log_buf.lock().unwrap().clone()).unwrap();
690+
assert!(
691+
output.contains("trace-test-id-12345"),
692+
"trace output should contain the request_id recorded in the span, got: {output}"
693+
);
694+
}
476695

477696
#[test]
478697
fn grpc_method_extracts_last_segment() {

0 commit comments

Comments
 (0)