77//! to either the gRPC service or HTTP endpoints based on the request headers.
88
99use bytes:: Bytes ;
10- use http:: { Request , Response } ;
10+ use http:: { HeaderValue , Request , Response } ;
1111use http_body:: Body ;
1212use http_body_util:: BodyExt ;
1313use hyper:: body:: Incoming ;
@@ -25,15 +25,86 @@ use std::sync::Arc;
2525use std:: task:: { Context , Poll } ;
2626use std:: time:: { Duration , Instant } ;
2727use 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 } ;
3030use tracing:: Span ;
3131
3232use 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-
428460fn 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) ]
474506mod 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