Skip to main content

toolkit_http/layers/
metrics.rs

1//! Tower layer that records OpenTelemetry metrics for outbound HTTP requests.
2//!
3//! Emits a single instrument following [OpenTelemetry HTTP client semantic
4//! conventions][semconv]:
5//! - `http.client.request.duration` — histogram (seconds)
6//!
7//! Attributes: `http.request.method`, `http.route`, `server.address`,
8//! `server.port` (when the URI carries an explicit port), and
9//! `http.response.status_code` (on success) or `error.type` (on failure).
10//!
11//! Modeled after go-appkit's `MetricsRoundTripper`: one duration histogram plus
12//! a build-time request classifier that produces the bounded `http.route`
13//! label, preventing cardinality explosion from raw paths. Like the Go version,
14//! this layer sits outside the retry loop, so it observes one logical request
15//! regardless of transport-level retries.
16//!
17//! [semconv]: https://opentelemetry.io/docs/specs/semconv/http/http-metrics/
18
19use crate::error::HttpError;
20use crate::request::RequestType;
21use crate::response::ResponseBody;
22use bytes::Bytes;
23use http::{Request, Response};
24use http_body_util::Full;
25use opentelemetry::metrics::{Histogram, Meter};
26use opentelemetry::{KeyValue, global};
27use std::borrow::Cow;
28use std::future::Future;
29use std::pin::Pin;
30use std::sync::Arc;
31use std::task::{Context, Poll};
32use std::time::Instant;
33use tower::{Layer, Service};
34
35/// Classifies a request into a low-cardinality route label (the `http.route`
36/// attribute). Set once when the client is built; invoked on every request.
37///
38/// This is the Rust analogue of go-appkit's `ClassifyRequest` callback. It must
39/// return a *bounded* set of values (e.g. route templates like
40/// `GET /users/{id}`), never a raw path containing identifiers, otherwise the
41/// metric cardinality is unbounded.
42pub type ClassifyFn = Arc<dyn Fn(&Request<Full<Bytes>>) -> Cow<'static, str> + Send + Sync>;
43
44/// Explicit histogram bucket boundaries (seconds) for request duration.
45///
46/// The SDK's default boundaries are count-oriented (hundreds–thousands) and
47/// useless for a seconds-valued duration. These mirror go-appkit's buckets with
48/// finer low-end resolution, so client-side percentiles stay meaningful and
49/// comparable across the two implementations.
50const DURATION_BOUNDARIES_SECS: &[f64] = &[
51    0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0, 60.0, 150.0, 300.0, 600.0,
52];
53
54/// Default classifier producing `"METHOD host"` (mirrors go-appkit's default
55/// `summary`). Never returns a raw path, so it cannot blow up cardinality.
56#[must_use]
57pub fn default_classify(req: &Request<Full<Bytes>>) -> Cow<'static, str> {
58    let host = req.uri().host().unwrap_or("unknown");
59    // Use the normalized method (`_OTHER` for unknown verbs) so the route label
60    // stays consistent with the `http.request.method` attribute and cannot be
61    // widened by arbitrary method strings.
62    Cow::Owned(format!("{} {}", normalize_method(req.method()), host))
63}
64
65/// Normalize HTTP method per [OTel semantic conventions][semconv].
66///
67/// Unknown methods map to `_OTHER` to bound attribute cardinality. Mirrors the
68/// server-side helper in `api-gateway`'s `http_metrics` middleware (duplicated
69/// here so `toolkit-http` stays free of a dependency on that gear).
70///
71/// [semconv]: https://opentelemetry.io/docs/specs/semconv/http/http-metrics/
72fn normalize_method(method: &http::Method) -> &'static str {
73    match *method {
74        http::Method::GET => "GET",
75        http::Method::POST => "POST",
76        http::Method::PUT => "PUT",
77        http::Method::DELETE => "DELETE",
78        http::Method::PATCH => "PATCH",
79        http::Method::HEAD => "HEAD",
80        http::Method::OPTIONS => "OPTIONS",
81        http::Method::CONNECT => "CONNECT",
82        http::Method::TRACE => "TRACE",
83        _ => "_OTHER",
84    }
85}
86
87/// Low-cardinality `error.type` value for a transport-level failure.
88///
89/// This layer sits outside the load-shed/retry layers (and inside the buffer),
90/// and the inner service returns `Ok(Response)` for all HTTP statuses (including
91/// 4xx/5xx). Only transport-class failures reach the `Err` arm here — the
92/// `OTel` analogue of go-appkit's `status="0"`. Because the layer is outside the
93/// concurrency limiter, a load-shed rejection reaches this arm too and is
94/// recorded as `"overloaded"`. Everything else collapses to `"other"` rather
95/// than enumerating variants that cannot occur at this point.
96fn error_type(err: &HttpError) -> &'static str {
97    match err {
98        HttpError::Timeout(_) => "timeout",
99        HttpError::DeadlineExceeded(_) => "deadline_exceeded",
100        HttpError::Transport(_) => "transport",
101        HttpError::Tls(_) => "tls",
102        HttpError::Overloaded => "overloaded",
103        _ => "other",
104    }
105}
106
107/// Tower layer recording HTTP client request-duration metrics.
108#[derive(Clone)]
109pub struct MetricsLayer {
110    duration: Histogram<f64>,
111    classify: ClassifyFn,
112}
113
114impl MetricsLayer {
115    /// Create a metrics layer.
116    ///
117    /// `client_type` names the OpenTelemetry instrumentation scope (the meter),
118    /// mirroring go-appkit's `ClientType` and the server-side `gear_name`.
119    /// `classify` produces the bounded `http.route` attribute for each request.
120    #[must_use]
121    pub fn new(client_type: &str, classify: ClassifyFn) -> Self {
122        let scope = opentelemetry::InstrumentationScope::builder(client_type.to_owned()).build();
123        let meter = global::meter_with_scope(scope);
124        Self::with_meter(&meter, classify)
125    }
126
127    /// Create a metrics layer using a caller-provided [`Meter`].
128    ///
129    /// Use this to bind the instrument to a specific `MeterProvider` instead of
130    /// the global one (e.g. for tests or multi-provider setups). The instrument
131    /// name, unit, bucket boundaries, and behavior are identical to [`new`](Self::new).
132    #[must_use]
133    pub fn with_meter(meter: &Meter, classify: ClassifyFn) -> Self {
134        let duration = meter
135            .f64_histogram("http.client.request.duration")
136            .with_description("Duration of outbound HTTP client requests")
137            .with_unit("s")
138            .with_boundaries(DURATION_BOUNDARIES_SECS.to_vec())
139            .build();
140        Self { duration, classify }
141    }
142}
143
144impl<S> Layer<S> for MetricsLayer {
145    type Service = MetricsService<S>;
146
147    fn layer(&self, inner: S) -> Self::Service {
148        MetricsService {
149            inner,
150            duration: self.duration.clone(),
151            classify: self.classify.clone(),
152        }
153    }
154}
155
156/// Service that records a duration metric for each outbound request.
157#[derive(Clone)]
158pub struct MetricsService<S> {
159    inner: S,
160    duration: Histogram<f64>,
161    classify: ClassifyFn,
162}
163
164impl<S> Service<Request<Full<Bytes>>> for MetricsService<S>
165where
166    S: Service<Request<Full<Bytes>>, Response = Response<ResponseBody>, Error = HttpError>
167        + Clone
168        + Send
169        + 'static,
170    S::Future: Send,
171{
172    type Response = Response<ResponseBody>;
173    type Error = HttpError;
174    type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
175
176    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
177        self.inner.poll_ready(cx)
178    }
179
180    fn call(&mut self, req: Request<Full<Bytes>>) -> Self::Future {
181        // Compute attributes from the request up front; the request itself is
182        // moved into the inner service.
183        let route = (self.classify)(&req).into_owned();
184        let method = normalize_method(req.method());
185        let server_address = req.uri().host().unwrap_or("unknown").to_owned();
186        let server_port = req.uri().port_u16();
187        // Read request_type set by RequestBuilder::with_request_type — mirrors
188        // go-appkit's GetRequestTypeFromContext.
189        let request_type = req
190            .extensions()
191            .get::<RequestType>()
192            .map(|rt| rt.0.clone().into_owned());
193        let duration = self.duration.clone();
194
195        // Swap so we call the instance that was poll_ready'd, leaving a fresh
196        // clone for the next poll_ready cycle (Tower Service contract).
197        let clone = self.inner.clone();
198        let mut inner = std::mem::replace(&mut self.inner, clone);
199
200        Box::pin(async move {
201            let start = Instant::now();
202            let result = inner.call(req).await;
203            let elapsed = start.elapsed().as_secs_f64();
204
205            let mut attrs = vec![
206                KeyValue::new("http.request.method", method),
207                KeyValue::new("http.route", route),
208                KeyValue::new("server.address", server_address),
209            ];
210            // OTel client semconv pairs server.address with server.port; only
211            // explicit ports are present in the URI (default 80/443 are elided).
212            if let Some(port) = server_port {
213                attrs.push(KeyValue::new("server.port", i64::from(port)));
214            }
215            if let Some(rt) = request_type {
216                attrs.push(KeyValue::new("request_type", rt));
217            }
218            match &result {
219                Ok(response) => attrs.push(KeyValue::new(
220                    "http.response.status_code",
221                    i64::from(response.status().as_u16()),
222                )),
223                Err(e) => attrs.push(KeyValue::new("error.type", error_type(e))),
224            }
225            duration.record(elapsed, &attrs);
226
227            result
228        })
229    }
230}
231
232#[cfg(test)]
233#[cfg_attr(coverage_nightly, coverage(off))]
234mod tests {
235    use super::*;
236    use crate::request::RequestType;
237    use http::StatusCode;
238    use http_body_util::{BodyExt, Empty};
239    use opentelemetry::metrics::MeterProvider;
240    use opentelemetry_sdk::metrics::data::{AggregatedMetrics, HistogramDataPoint, MetricData};
241    use opentelemetry_sdk::metrics::{InMemoryMetricExporter, SdkMeterProvider};
242    use std::convert::Infallible;
243    use tower::{ServiceBuilder, ServiceExt, service_fn};
244
245    fn empty_response(status: StatusCode) -> Response<ResponseBody> {
246        let body: ResponseBody = Empty::<Bytes>::new()
247            .map_err(|e: Infallible| -> Box<dyn std::error::Error + Send + Sync> { match e {} })
248            .boxed();
249        Response::builder().status(status).body(body).unwrap()
250    }
251
252    /// Collect the histogram data point for `http.client.request.duration` whose
253    /// attributes contain every `(key, value)` in `expected`. Returns `None` if
254    /// no matching point was exported.
255    fn find_duration_point(
256        exporter: &InMemoryMetricExporter,
257        expected: &[(&str, &str)],
258    ) -> Option<HistogramDataPoint<f64>> {
259        let batches = exporter.get_finished_metrics().unwrap();
260        for rm in &batches {
261            for sm in rm.scope_metrics() {
262                for metric in sm.metrics() {
263                    if metric.name() != "http.client.request.duration" {
264                        continue;
265                    }
266                    let AggregatedMetrics::F64(MetricData::Histogram(hist)) = metric.data() else {
267                        continue;
268                    };
269                    for dp in hist.data_points() {
270                        let matches = expected.iter().all(|(k, v)| {
271                            dp.attributes()
272                                .any(|kv| kv.key.as_str() == *k && kv.value.to_string() == *v)
273                        });
274                        if matches {
275                            return Some(dp.clone());
276                        }
277                    }
278                }
279            }
280        }
281        None
282    }
283
284    fn test_provider() -> (SdkMeterProvider, InMemoryMetricExporter) {
285        let exporter = InMemoryMetricExporter::default();
286        let provider = SdkMeterProvider::builder()
287            .with_periodic_exporter(exporter.clone())
288            .build();
289        (provider, exporter)
290    }
291
292    #[tokio::test]
293    async fn records_duration_with_attributes_on_success() {
294        let (provider, exporter) = test_provider();
295        let meter = provider.meter("test-client");
296        let classify: ClassifyFn = Arc::new(|_req| Cow::Borrowed("GET /users/{id}"));
297        let layer = MetricsLayer::with_meter(&meter, classify);
298
299        let inner = service_fn(|_req: Request<Full<Bytes>>| async {
300            Ok::<_, HttpError>(empty_response(StatusCode::OK))
301        });
302        let mut svc = ServiceBuilder::new().layer(layer).service(inner);
303        let req = Request::builder()
304            .method(http::Method::GET)
305            .uri("https://example.com:8443/users/123")
306            .body(Full::new(Bytes::new()))
307            .unwrap();
308
309        let resp = svc.ready().await.unwrap().call(req).await.unwrap();
310        assert_eq!(resp.status(), StatusCode::OK);
311
312        provider.force_flush().unwrap();
313        let point = find_duration_point(
314            &exporter,
315            &[
316                ("http.request.method", "GET"),
317                ("http.route", "GET /users/{id}"),
318                ("server.address", "example.com"),
319                ("server.port", "8443"),
320                ("http.response.status_code", "200"),
321            ],
322        )
323        .expect("a duration data point with the expected attributes should be exported");
324        assert_eq!(point.count(), 1, "exactly one observation recorded");
325    }
326
327    #[tokio::test]
328    async fn records_error_type_on_transport_failure() {
329        let (provider, exporter) = test_provider();
330        let meter = provider.meter("test-client");
331        let layer = MetricsLayer::with_meter(&meter, Arc::new(default_classify));
332
333        let inner = service_fn(|_req: Request<Full<Bytes>>| async {
334            Err::<Response<ResponseBody>, _>(HttpError::Timeout(std::time::Duration::from_secs(1)))
335        });
336        let mut svc = ServiceBuilder::new().layer(layer).service(inner);
337        let req = Request::builder()
338            .method(http::Method::GET)
339            .uri("https://example.com/")
340            .body(Full::new(Bytes::new()))
341            .unwrap();
342
343        let err = svc.ready().await.unwrap().call(req).await.unwrap_err();
344        assert!(matches!(err, HttpError::Timeout(_)));
345
346        provider.force_flush().unwrap();
347        let point = find_duration_point(&exporter, &[("error.type", "timeout")])
348            .expect("a duration data point tagged error.type=timeout should be exported");
349        assert_eq!(point.count(), 1);
350        // Failures must not carry a status code.
351        assert!(
352            point
353                .attributes()
354                .all(|kv| kv.key.as_str() != "http.response.status_code"),
355            "transport failures must not record http.response.status_code"
356        );
357    }
358
359    #[test]
360    fn default_classify_normalizes_method_and_drops_path() {
361        let req = Request::builder()
362            .method(http::Method::POST)
363            .uri("https://api.example.com/users/abc-123-uuid")
364            .body(Full::new(Bytes::new()))
365            .unwrap();
366        // Raw path with an identifier must never leak into the label.
367        assert_eq!(default_classify(&req), "POST api.example.com");
368
369        let exotic = Request::builder()
370            .method(http::Method::from_bytes(b"PROPFIND").unwrap())
371            .uri("https://api.example.com/dav")
372            .body(Full::new(Bytes::new()))
373            .unwrap();
374        assert_eq!(default_classify(&exotic), "_OTHER api.example.com");
375    }
376
377    #[test]
378    fn normalize_method_caps_unknown() {
379        assert_eq!(normalize_method(&http::Method::GET), "GET");
380        let custom = http::Method::from_bytes(b"PROPFIND").unwrap();
381        assert_eq!(normalize_method(&custom), "_OTHER");
382    }
383
384    #[tokio::test]
385    async fn records_request_type_attribute_when_set() {
386        let (provider, exporter) = test_provider();
387        let meter = provider.meter("test-client");
388        let layer = MetricsLayer::with_meter(&meter, Arc::new(default_classify));
389
390        let inner = service_fn(|_req: Request<Full<Bytes>>| async {
391            Ok::<_, HttpError>(empty_response(StatusCode::OK))
392        });
393        let mut svc = ServiceBuilder::new().layer(layer).service(inner);
394
395        let mut req = Request::builder()
396            .method(http::Method::GET)
397            .uri("https://example.com/tenants/123")
398            .body(Full::new(Bytes::new()))
399            .unwrap();
400        req.extensions_mut()
401            .insert(RequestType::new("tenants_resolve"));
402
403        svc.ready().await.unwrap().call(req).await.unwrap();
404
405        provider.force_flush().unwrap();
406        let point = find_duration_point(&exporter, &[("request_type", "tenants_resolve")])
407            .expect("request_type attribute should appear in exported metric");
408        assert_eq!(point.count(), 1);
409    }
410
411    #[tokio::test]
412    async fn omits_request_type_when_not_set() {
413        let (provider, exporter) = test_provider();
414        let meter = provider.meter("test-client");
415        let layer = MetricsLayer::with_meter(&meter, Arc::new(default_classify));
416
417        let inner = service_fn(|_req: Request<Full<Bytes>>| async {
418            Ok::<_, HttpError>(empty_response(StatusCode::OK))
419        });
420        let mut svc = ServiceBuilder::new().layer(layer).service(inner);
421
422        let req = Request::builder()
423            .method(http::Method::GET)
424            .uri("https://example.com/tenants/123")
425            .body(Full::new(Bytes::new()))
426            .unwrap();
427
428        svc.ready().await.unwrap().call(req).await.unwrap();
429
430        provider.force_flush().unwrap();
431        let dp = find_duration_point(&exporter, &[("http.request.method", "GET")])
432            .expect("a data point should be exported");
433        assert!(
434            dp.attributes().all(|kv| kv.key.as_str() != "request_type"),
435            "request_type must not appear when not set"
436        );
437    }
438
439    #[test]
440    fn error_type_maps_transport_class_failures() {
441        assert_eq!(
442            error_type(&HttpError::Timeout(std::time::Duration::from_secs(1))),
443            "timeout"
444        );
445        assert_eq!(
446            error_type(&HttpError::Transport("boom".into())),
447            "transport"
448        );
449        assert_eq!(error_type(&HttpError::Overloaded), "overloaded");
450    }
451
452    /// A load-shed rejection (`HttpError::Overloaded`) reaching this layer — as
453    /// it does in `build()`, where the concurrency limiter is inner to the
454    /// metrics layer — is recorded as `error.type = "overloaded"`, so shed
455    /// requests are counted rather than invisible.
456    #[tokio::test]
457    async fn records_error_type_overloaded_when_shed() {
458        let (provider, exporter) = test_provider();
459        let meter = provider.meter("test-client");
460        let layer = MetricsLayer::with_meter(&meter, Arc::new(default_classify));
461
462        let inner = service_fn(|_req: Request<Full<Bytes>>| async {
463            Err::<Response<ResponseBody>, _>(HttpError::Overloaded)
464        });
465        let mut svc = ServiceBuilder::new().layer(layer).service(inner);
466        let req = Request::builder()
467            .method(http::Method::GET)
468            .uri("https://example.com/")
469            .body(Full::new(Bytes::new()))
470            .unwrap();
471
472        let err = svc.ready().await.unwrap().call(req).await.unwrap_err();
473        assert!(matches!(err, HttpError::Overloaded));
474
475        provider.force_flush().unwrap();
476        let point = find_duration_point(&exporter, &[("error.type", "overloaded")])
477            .expect("a shed request should record error.type=overloaded");
478        assert_eq!(point.count(), 1);
479    }
480}