Skip to main content

api_bones_tower/
lib.rs

1//! Tower middleware building blocks for api-bones services.
2//!
3//! Provides composable Tower [`Layer`](tower::Layer) / [`Service`](tower::Service)
4//! implementations for:
5//!
6//! | Layer                 | What it does                                         |
7//! |-----------------------|------------------------------------------------------|
8//! | [`RequestIdLayer`]    | Generates / propagates `X-Request-Id` on every req  |
9//! | [`ProblemJsonLayer`]  | Maps non-`ApiError` inner-service errors to Problem+JSON |
10//! | `TraceContextLayer`   | Injects W3C `traceparent`/`tracestate` on every outbound request (feature `opentelemetry`) |
11//!
12//! ## Feature flags
13//!
14//! By default this crate enables `std` and `serde` on `api-bones`.
15//! Additional `api-bones` features can be opted into:
16//!
17//! | Feature         | What it enables                                                   |
18//! |-----------------|-------------------------------------------------------------------|
19//! | `uuid`          | UUID-based request IDs (`api-bones/uuid`)                         |
20//! | `chrono`        | Chrono timestamp types (`api-bones/chrono`)                       |
21//! | `opentelemetry` | Client-side `TraceContextLayer` — outbound W3C trace-context injection |
22//!
23//! # Example
24//!
25//! ```rust,no_run
26//! use api_bones_tower::{RequestIdLayer, ProblemJsonLayer};
27//! use tower::ServiceBuilder;
28//!
29//! let _svc = ServiceBuilder::new()
30//!     .layer(RequestIdLayer::new())
31//!     .layer(ProblemJsonLayer::new())
32//!     .service(tower::service_fn(|_req: http::Request<()>| async {
33//!         Ok::<_, std::convert::Infallible>(http::Response::new(()))
34//!     }));
35//! ```
36
37use std::future::Future;
38use std::pin::Pin;
39use std::sync::Arc;
40use std::sync::atomic::{AtomicU64, Ordering};
41use std::task::{Context, Poll};
42
43use api_bones::error::ApiError;
44use http::{Request, Response};
45use tower::{Layer, Service};
46
47// ---------------------------------------------------------------------------
48// RequestIdLayer
49// ---------------------------------------------------------------------------
50
51/// Tower [`Layer`] that ensures every request carries an `X-Request-Id` header.
52///
53/// - If the incoming request already has an `X-Request-Id`, it is forwarded
54///   unchanged.
55/// - Otherwise a monotonically-increasing numeric ID is generated and injected
56///   (format: `req-<n>`).
57///
58/// The same header value is echoed back in the response.
59///
60/// # Example
61///
62/// ```rust,no_run
63/// use api_bones_tower::RequestIdLayer;
64/// use tower::ServiceBuilder;
65///
66/// let _svc = ServiceBuilder::new()
67///     .layer(RequestIdLayer::new())
68///     .service(tower::service_fn(|_req: http::Request<()>| async {
69///         Ok::<_, std::convert::Infallible>(http::Response::new(()))
70///     }));
71/// ```
72#[derive(Clone, Debug)]
73pub struct RequestIdLayer {
74    counter: Arc<AtomicU64>,
75}
76
77impl RequestIdLayer {
78    /// Create a new `RequestIdLayer` with an internal counter starting at 1.
79    #[must_use]
80    pub fn new() -> Self {
81        Self {
82            counter: Arc::new(AtomicU64::new(1)),
83        }
84    }
85}
86
87impl Default for RequestIdLayer {
88    fn default() -> Self {
89        Self::new()
90    }
91}
92
93impl<S> Layer<S> for RequestIdLayer {
94    type Service = RequestIdService<S>;
95
96    fn layer(&self, inner: S) -> Self::Service {
97        RequestIdService {
98            inner,
99            counter: Arc::clone(&self.counter),
100        }
101    }
102}
103
104/// Tower [`Service`] produced by [`RequestIdLayer`].
105#[derive(Clone, Debug)]
106pub struct RequestIdService<S> {
107    inner: S,
108    counter: Arc<AtomicU64>,
109}
110
111impl<S, ReqBody, ResBody> Service<Request<ReqBody>> for RequestIdService<S>
112where
113    S: Service<Request<ReqBody>, Response = Response<ResBody>> + Clone + Send + 'static,
114    S::Future: Send,
115    S::Error: Send,
116    ReqBody: Send + 'static,
117    ResBody: Default + Send,
118{
119    type Response = Response<ResBody>;
120    type Error = S::Error;
121    type Future = RequestIdFuture<S::Future, ResBody>;
122
123    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
124        self.inner.poll_ready(cx)
125    }
126
127    fn call(&mut self, mut req: Request<ReqBody>) -> Self::Future {
128        // Determine (or generate) the request ID.
129        let request_id: String = if let Some(existing) = req.headers().get("x-request-id") {
130            existing.to_str().unwrap_or("invalid").to_owned()
131        } else {
132            let n = self.counter.fetch_add(1, Ordering::Relaxed);
133            let id = format!("req-{n}");
134            if let Ok(val) = http::HeaderValue::from_str(&id) {
135                req.headers_mut().insert("x-request-id", val);
136            }
137            id
138        };
139
140        let future = self.inner.call(req);
141        RequestIdFuture {
142            inner: future,
143            request_id,
144            _body: std::marker::PhantomData,
145        }
146    }
147}
148
149/// Future returned by [`RequestIdService`].
150#[pin_project::pin_project]
151pub struct RequestIdFuture<F, ResBody> {
152    #[pin]
153    inner: F,
154    request_id: String,
155    _body: std::marker::PhantomData<ResBody>,
156}
157
158impl<F, ResBody, E> Future for RequestIdFuture<F, ResBody>
159where
160    F: Future<Output = Result<Response<ResBody>, E>>,
161{
162    type Output = Result<Response<ResBody>, E>;
163
164    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
165        let this = self.project();
166        match this.inner.poll(cx) {
167            Poll::Pending => Poll::Pending,
168            Poll::Ready(Ok(mut resp)) => {
169                if let Ok(val) = http::HeaderValue::from_str(this.request_id) {
170                    resp.headers_mut().entry("x-request-id").or_insert(val);
171                }
172                Poll::Ready(Ok(resp))
173            }
174            Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
175        }
176    }
177}
178
179// ---------------------------------------------------------------------------
180// ProblemJsonLayer
181// ---------------------------------------------------------------------------
182
183/// Tower [`Layer`] that maps inner-service errors into Problem+JSON HTTP
184/// responses.
185///
186/// Any `Err` propagated from the inner service is converted to an [`ApiError`]
187/// via the [`Into<ApiError>`] bound and then serialized as
188/// `application/problem+json`.
189///
190/// Successful responses are passed through unchanged.
191///
192/// # Example
193///
194/// ```rust,no_run
195/// use api_bones_tower::ProblemJsonLayer;
196/// use tower::ServiceBuilder;
197///
198/// let _svc = ServiceBuilder::new()
199///     .layer(ProblemJsonLayer::new())
200///     .service(tower::service_fn(|_req: http::Request<()>| async {
201///         Ok::<_, api_bones::ApiError>(http::Response::new(String::new()))
202///     }));
203/// ```
204#[derive(Clone, Debug, Default)]
205pub struct ProblemJsonLayer;
206
207impl ProblemJsonLayer {
208    /// Create a new `ProblemJsonLayer`.
209    #[must_use]
210    pub fn new() -> Self {
211        Self
212    }
213}
214
215impl<S> Layer<S> for ProblemJsonLayer {
216    type Service = ProblemJsonService<S>;
217
218    fn layer(&self, inner: S) -> Self::Service {
219        ProblemJsonService { inner }
220    }
221}
222
223/// Tower [`Service`] produced by [`ProblemJsonLayer`].
224#[derive(Clone, Debug)]
225pub struct ProblemJsonService<S> {
226    inner: S,
227}
228
229impl<S, ReqBody> Service<Request<ReqBody>> for ProblemJsonService<S>
230where
231    S: Service<Request<ReqBody>, Response = Response<String>> + Clone + Send + 'static,
232    S::Error: Into<ApiError> + Send,
233    S::Future: Send,
234    ReqBody: Send + 'static,
235{
236    type Response = Response<String>;
237    type Error = std::convert::Infallible;
238    type Future =
239        Pin<Box<dyn Future<Output = Result<Response<String>, std::convert::Infallible>> + Send>>;
240
241    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
242        match self.inner.poll_ready(cx) {
243            Poll::Pending => Poll::Pending,
244            Poll::Ready(Ok(())) => Poll::Ready(Ok(())),
245            Poll::Ready(Err(_e)) => unreachable!("inner service poll_ready returned Err"),
246        }
247    }
248
249    fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
250        let future = self.inner.call(req);
251        Box::pin(async move {
252            match future.await {
253                Ok(resp) => Ok(resp),
254                Err(e) => {
255                    let api_err: ApiError = e.into();
256                    Ok(api_error_to_response(api_err))
257                }
258            }
259        })
260    }
261}
262
263/// Convert an [`ApiError`] into an HTTP response with `application/problem+json`.
264fn api_error_to_response(err: ApiError) -> Response<String> {
265    use api_bones::error::ProblemJson;
266
267    let status = err.status;
268    let problem = ProblemJson::from(err);
269    let body = serde_json::to_string(&problem).expect("ProblemJson serialization is infallible");
270
271    let status_code =
272        http::StatusCode::from_u16(status).unwrap_or(http::StatusCode::INTERNAL_SERVER_ERROR);
273
274    Response::builder()
275        .status(status_code)
276        .header("content-type", "application/problem+json")
277        .body(body)
278        .expect("response construction is infallible for valid status codes")
279}
280
281// ---------------------------------------------------------------------------
282// TraceContextLayer
283// ---------------------------------------------------------------------------
284
285/// Tower [`Layer`] that injects the active OpenTelemetry trace context into
286/// every outbound request's headers.
287///
288/// It emits W3C `traceparent` / `tracestate` so the callee's span links to the
289/// caller's span instead of starting a fresh orphan-root trace.
290///
291/// This is the client-side complement of the server-side context extraction
292/// the inbound middleware performs. Stack it once on any [`tower::Service`]-based
293/// HTTP transport — for example a `connectrpc` client transport, which is a
294/// `tower::Service<http::Request<_>>` — and **every** call made through that
295/// client propagates trace context automatically: no per-call code, no per-SDK
296/// `inject_current` boilerplate. That "inject on every outbound call, for any
297/// reason" property is the whole point — a background poller's calls propagate
298/// exactly like a request handler's, as long as the caller runs inside a span.
299///
300/// Injection reads [`opentelemetry::Context::current`] at dispatch time, so it
301/// carries whatever span is active on the calling task. When no span is active
302/// the propagator emits nothing, so the layer is a transparent no-op on
303/// untraced calls rather than a source of malformed headers.
304///
305/// # Example
306///
307/// ```rust,no_run
308/// use api_bones_tower::TraceContextLayer;
309/// use tower::ServiceBuilder;
310///
311/// let _svc = ServiceBuilder::new()
312///     .layer(TraceContextLayer::new())
313///     .service(tower::service_fn(|_req: http::Request<()>| async {
314///         Ok::<_, std::convert::Infallible>(http::Response::new(()))
315///     }));
316/// ```
317#[cfg(feature = "opentelemetry")]
318#[derive(Clone, Debug, Default)]
319pub struct TraceContextLayer;
320
321#[cfg(feature = "opentelemetry")]
322impl TraceContextLayer {
323    /// Create a new `TraceContextLayer`.
324    #[must_use]
325    pub fn new() -> Self {
326        Self
327    }
328}
329
330#[cfg(feature = "opentelemetry")]
331impl<S> Layer<S> for TraceContextLayer {
332    type Service = TraceContextService<S>;
333
334    fn layer(&self, inner: S) -> Self::Service {
335        TraceContextService { inner }
336    }
337}
338
339/// Tower [`Service`] produced by [`TraceContextLayer`].
340#[cfg(feature = "opentelemetry")]
341#[derive(Clone, Debug)]
342pub struct TraceContextService<S> {
343    inner: S,
344}
345
346#[cfg(feature = "opentelemetry")]
347impl<S, ReqBody, ResBody> Service<Request<ReqBody>> for TraceContextService<S>
348where
349    S: Service<Request<ReqBody>, Response = Response<ResBody>>,
350{
351    type Response = Response<ResBody>;
352    type Error = S::Error;
353    type Future = S::Future;
354
355    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
356        self.inner.poll_ready(cx)
357    }
358
359    fn call(&mut self, mut req: Request<ReqBody>) -> Self::Future {
360        // Inject the active W3C trace context into the outbound request headers
361        // so the callee's span links to the caller's. A no-op when no span is
362        // active on the current task.
363        api_bones::propagation::inject_current(req.headers_mut());
364        self.inner.call(req)
365    }
366}
367
368// ---------------------------------------------------------------------------
369// Tests
370// ---------------------------------------------------------------------------
371
372#[cfg(test)]
373mod tests {
374    use super::*;
375    use tower::{ServiceBuilder, ServiceExt};
376
377    #[tokio::test]
378    async fn request_id_layer_injects_header() {
379        let svc = ServiceBuilder::new()
380            .layer(RequestIdLayer::new())
381            .service(tower::service_fn(|req: Request<()>| async move {
382                let id = req
383                    .headers()
384                    .get("x-request-id")
385                    .and_then(|v| v.to_str().ok())
386                    .unwrap_or("")
387                    .to_owned();
388                let resp = Response::new(id);
389                Ok::<_, std::convert::Infallible>(resp)
390            }));
391
392        let req = Request::builder().uri("/").body(()).unwrap();
393        let resp = svc.oneshot(req).await.unwrap();
394        assert!(resp.headers().contains_key("x-request-id"));
395    }
396
397    #[tokio::test]
398    async fn request_id_layer_preserves_existing_header() {
399        let svc = ServiceBuilder::new()
400            .layer(RequestIdLayer::new())
401            .service(tower::service_fn(|_req: Request<()>| async move {
402                Ok::<_, std::convert::Infallible>(Response::new(String::new()))
403            }));
404
405        let req = Request::builder()
406            .uri("/")
407            .header("x-request-id", "client-id")
408            .body(())
409            .unwrap();
410        let resp = svc.oneshot(req).await.unwrap();
411        assert_eq!(
412            resp.headers()
413                .get("x-request-id")
414                .unwrap()
415                .to_str()
416                .unwrap(),
417            "client-id"
418        );
419    }
420
421    #[tokio::test]
422    async fn problem_json_layer_maps_error() {
423        let svc = ServiceBuilder::new()
424            .layer(ProblemJsonLayer::new())
425            .service(tower::service_fn(|_req: Request<()>| async move {
426                Err::<Response<String>, ApiError>(ApiError::not_found("item 1"))
427            }));
428
429        let req = Request::builder().uri("/").body(()).unwrap();
430        let resp = svc.oneshot(req).await.unwrap();
431        assert_eq!(resp.status().as_u16(), 404);
432        assert_eq!(
433            resp.headers()
434                .get("content-type")
435                .unwrap()
436                .to_str()
437                .unwrap(),
438            "application/problem+json"
439        );
440    }
441
442    #[tokio::test]
443    async fn problem_json_layer_passes_through_ok() {
444        let svc = ServiceBuilder::new()
445            .layer(ProblemJsonLayer::new())
446            .service(tower::service_fn(|_req: Request<()>| async move {
447                Ok::<_, ApiError>(
448                    Response::builder()
449                        .status(200)
450                        .body("ok".to_owned())
451                        .unwrap(),
452                )
453            }));
454
455        let req = Request::builder().uri("/").body(()).unwrap();
456        let resp = svc.oneshot(req).await.unwrap();
457        assert_eq!(resp.status().as_u16(), 200);
458    }
459
460    #[test]
461    fn request_id_layer_default_is_same_as_new() {
462        let _layer = RequestIdLayer::default();
463    }
464
465    #[tokio::test]
466    async fn problem_json_service_poll_ready() {
467        use tower::{Service, ServiceExt};
468
469        let inner = tower::service_fn(|_req: Request<()>| async move {
470            Ok::<_, ApiError>(Response::builder().body("ok".to_owned()).unwrap())
471        });
472        let mut svc = ProblemJsonService { inner };
473        let svc_ref = svc.ready().await.unwrap();
474        let req = Request::builder().uri("/").body(()).unwrap();
475        let resp = svc_ref.call(req).await.unwrap();
476        assert_eq!(resp.status().as_u16(), 200);
477    }
478
479    #[tokio::test]
480    async fn request_id_future_propagates_inner_error() {
481        let svc = ServiceBuilder::new()
482            .layer(RequestIdLayer::new())
483            .service(tower::service_fn(|_req: Request<()>| async move {
484                Err::<Response<String>, ApiError>(ApiError::internal("boom"))
485            }));
486
487        let req = Request::builder().uri("/").body(()).unwrap();
488        let result = svc.oneshot(req).await;
489        let err = result.unwrap_err();
490        assert_eq!(err.status, 500);
491    }
492
493    #[tokio::test]
494    async fn request_id_future_poll_pending() {
495        use std::sync::{
496            Arc,
497            atomic::{AtomicBool, Ordering},
498        };
499
500        let ready = Arc::new(AtomicBool::new(false));
501        let ready2 = Arc::clone(&ready);
502
503        let inner = tower::service_fn(move |_req: Request<()>| {
504            let flag = Arc::clone(&ready2);
505            async move {
506                tokio::task::yield_now().await;
507                flag.store(true, Ordering::SeqCst);
508                Ok::<Response<String>, std::convert::Infallible>(
509                    Response::builder().body(String::new()).unwrap(),
510                )
511            }
512        });
513
514        let layer = RequestIdLayer::new();
515        let mut svc = layer.layer(inner);
516
517        let req = Request::builder().uri("/").body(()).unwrap();
518        let fut = tower::Service::call(&mut svc, req);
519        let resp = fut.await.unwrap();
520        assert!(resp.headers().contains_key("x-request-id"));
521        assert!(ready.load(Ordering::SeqCst));
522    }
523
524    #[cfg(feature = "opentelemetry")]
525    #[tokio::test]
526    async fn trace_context_layer_injects_traceparent_on_outbound() {
527        use std::sync::{Arc, Mutex};
528
529        use opentelemetry::Context as OtelContext;
530        use opentelemetry::global;
531        use opentelemetry::trace::{TraceContextExt as _, Tracer as _, TracerProvider as _};
532        use opentelemetry_sdk::propagation::TraceContextPropagator;
533        use opentelemetry_sdk::trace::TracerProvider as SdkTracerProvider;
534        use tower::ServiceExt as _;
535
536        global::set_text_map_propagator(TraceContextPropagator::new());
537        let provider = SdkTracerProvider::builder().build();
538        let tracer = provider.tracer("test");
539        let span = tracer.start("caller-span");
540        let _guard = OtelContext::current_with_span(span).attach();
541
542        // Capture the headers the inner (transport) service actually receives.
543        let seen: Arc<Mutex<Option<http::HeaderMap>>> = Arc::new(Mutex::new(None));
544        let seen_inner = Arc::clone(&seen);
545        let inner = tower::service_fn(move |req: Request<()>| {
546            let seen = Arc::clone(&seen_inner);
547            async move {
548                *seen.lock().expect("headers mutex poisoned") = Some(req.headers().clone());
549                Ok::<Response<()>, std::convert::Infallible>(Response::new(()))
550            }
551        });
552
553        let svc = TraceContextLayer::new().layer(inner);
554        let req = Request::builder()
555            .uri("/")
556            .body(())
557            .expect("request builds");
558        svc.oneshot(req).await.expect("service call succeeds");
559
560        let headers = seen
561            .lock()
562            .expect("headers mutex poisoned")
563            .take()
564            .expect("inner service ran");
565        assert!(
566            headers.contains_key("traceparent"),
567            "expected traceparent injected into outbound headers, got: {headers:?}"
568        );
569    }
570
571    #[cfg(feature = "opentelemetry")]
572    #[tokio::test]
573    async fn trace_context_layer_is_noop_without_active_span() {
574        use std::sync::{Arc, Mutex};
575
576        use tower::ServiceExt as _;
577
578        let seen: Arc<Mutex<Option<http::HeaderMap>>> = Arc::new(Mutex::new(None));
579        let seen_inner = Arc::clone(&seen);
580        let inner = tower::service_fn(move |req: Request<()>| {
581            let seen = Arc::clone(&seen_inner);
582            async move {
583                *seen.lock().expect("headers mutex poisoned") = Some(req.headers().clone());
584                Ok::<Response<()>, std::convert::Infallible>(Response::new(()))
585            }
586        });
587
588        let svc = TraceContextLayer::new().layer(inner);
589        let req = Request::builder()
590            .uri("/")
591            .body(())
592            .expect("request builds");
593        svc.oneshot(req).await.expect("service call succeeds");
594
595        let headers = seen
596            .lock()
597            .expect("headers mutex poisoned")
598            .take()
599            .expect("inner service ran");
600        assert!(
601            !headers.contains_key("traceparent"),
602            "expected no traceparent header when no span is active, got: {headers:?}"
603        );
604    }
605}