1use 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
35pub type ClassifyFn = Arc<dyn Fn(&Request<Full<Bytes>>) -> Cow<'static, str> + Send + Sync>;
43
44const 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#[must_use]
57pub fn default_classify(req: &Request<Full<Bytes>>) -> Cow<'static, str> {
58 let host = req.uri().host().unwrap_or("unknown");
59 Cow::Owned(format!("{} {}", normalize_method(req.method()), host))
63}
64
65fn 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
87fn 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#[derive(Clone)]
109pub struct MetricsLayer {
110 duration: Histogram<f64>,
111 classify: ClassifyFn,
112}
113
114impl MetricsLayer {
115 #[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 #[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#[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 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 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 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 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 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 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 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 #[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}