1use 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#[derive(Clone, Debug)]
73pub struct RequestIdLayer {
74 counter: Arc<AtomicU64>,
75}
76
77impl RequestIdLayer {
78 #[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#[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 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#[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#[derive(Clone, Debug, Default)]
205pub struct ProblemJsonLayer;
206
207impl ProblemJsonLayer {
208 #[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#[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
263fn 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#[cfg(feature = "opentelemetry")]
318#[derive(Clone, Debug, Default)]
319pub struct TraceContextLayer;
320
321#[cfg(feature = "opentelemetry")]
322impl TraceContextLayer {
323 #[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#[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 api_bones::propagation::inject_current(req.headers_mut());
364 self.inner.call(req)
365 }
366}
367
368#[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 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}