use std::{ collections::HashMap, net::SocketAddr, sync::{ Arc, atomic::{AtomicUsize, Ordering}, }, }; use axum::{ Router, extract::{Request, State}, http::{ HeaderMap, StatusCode, header::{self, HeaderValue}, }, middleware::{self, Next}, response::{IntoResponse, Response}, routing::get, }; use metrics_exporter_prometheus::{PrometheusBuilder, PrometheusHandle, PrometheusRecorder}; use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; use thiserror::Error; use tokio::net::TcpListener; use crate::{DURATION_BUCKETS_SECONDS, ServiceIdentity}; use crank_metrics::{ExemplarObservation, MAX_EXPOSITION_BYTES, MetricService, exemplar_snapshot}; const PROMETHEUS_CONTENT_TYPE: &str = "text/plain; version=0.0.4; charset=utf-8"; const OPENMETRICS_CONTENT_TYPE: &str = "application/openmetrics-text; version=1.0.0; charset=utf-8"; const EXPOSITION_TOO_LARGE: &str = "metrics exposition exceeds configured bound\n"; const TOO_MANY_SCRAPES: &str = "metrics scrape concurrency limit exceeded\n"; const MAX_CONCURRENT_SCRAPES: usize = 2; #[derive(Clone)] pub struct MetricsConfig { enabled: bool, bind_addr: SocketAddr, token_digest: Option<[u8; 32]>, } impl std::fmt::Debug for MetricsConfig { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter .debug_struct("MetricsConfig") .field("enabled", &self.enabled) .field("bind_addr", &self.bind_addr) .field("authentication_configured", &self.token_digest.is_some()) .finish() } } impl MetricsConfig { pub fn new( enabled: bool, bind_addr: SocketAddr, bearer_token: Option, ) -> Result { let token_digest = bearer_token .filter(|token| !token.is_empty()) .map(|token| token_digest(token.as_bytes())); if enabled && !bind_addr.ip().is_loopback() && token_digest.is_none() { return Err(MetricsConfigError::MissingTokenForExternalBind); } Ok(Self { enabled, bind_addr, token_digest, }) } pub fn enabled(&self) -> bool { self.enabled } pub fn bind_addr(&self) -> SocketAddr { self.bind_addr } pub fn requires_authentication(&self) -> bool { !self.bind_addr.ip().is_loopback() } } #[derive(Debug, Error)] pub enum MetricsConfigError { #[error("metrics environment variable is not valid UTF-8: {field}")] InvalidEnvironmentEncoding { field: &'static str }, #[error("metrics bind address is invalid: {field}")] InvalidBindAddress { field: &'static str }, #[error("metrics enabled flag must be one of true, false, 1, 0")] InvalidEnabledFlag, #[error("external metrics bind requires a bearer token")] MissingTokenForExternalBind, } #[derive(Clone)] struct MetricsState { handle: PrometheusHandle, token_digest: Option<[u8; 32]>, requires_authentication: bool, active_scrapes: Arc, } pub struct MetricsSurface { config: MetricsConfig, state: MetricsState, _recorder: Option, } impl MetricsSurface { pub(crate) fn new(config: MetricsConfig, handle: PrometheusHandle) -> Self { Self { state: MetricsState { handle, token_digest: config.token_digest, requires_authentication: config.requires_authentication(), active_scrapes: Arc::new(AtomicUsize::new(0)), }, config, _recorder: None, } } pub fn for_test( config: MetricsConfig, identity: ServiceIdentity, ) -> Result { let recorder = prometheus_builder(&identity)?.build_recorder(); let handle = recorder.handle(); let mut surface = Self::new(config, handle); surface._recorder = Some(recorder); Ok(surface) } pub fn router(&self) -> Router { Router::new() .route("/metrics", get(render_metrics)) .route("/health", get(metrics_health)) .route_layer(middleware::from_fn_with_state( self.state.clone(), authorize_metrics, )) .with_state(self.state.clone()) } pub async fn bind(self) -> Result { let listener = TcpListener::bind(self.config.bind_addr) .await .map_err(|_| MetricsServeError::Bind)?; Ok(MetricsServer { listener, router: self.router(), }) } } pub struct MetricsServer { listener: TcpListener, router: Router, } impl MetricsServer { pub fn local_addr(&self) -> Result { self.listener .local_addr() .map_err(|_| MetricsServeError::LocalAddress) } pub async fn serve(self) -> Result<(), MetricsServeError> { axum::serve(self.listener, self.router) .await .map_err(|_| MetricsServeError::Serve) } } #[derive(Debug, Error)] pub enum MetricsSurfaceError { #[error("failed to configure Prometheus recorder")] RecorderConfiguration, } #[derive(Debug, Error)] pub enum MetricsServeError { #[error("failed to bind metrics listener")] Bind, #[error("metrics listener stopped unexpectedly")] Serve, #[error("failed to read metrics listener address")] LocalAddress, } pub(crate) fn install_prometheus_recorder( identity: &ServiceIdentity, ) -> Result { prometheus_builder(identity)? .install_recorder() .map_err(|_| MetricsSurfaceError::RecorderConfiguration) } fn prometheus_builder( identity: &ServiceIdentity, ) -> Result { MetricService::parse(identity.service()).ok_or(MetricsSurfaceError::RecorderConfiguration)?; PrometheusBuilder::new() .set_buckets(DURATION_BUCKETS_SECONDS) .map(|builder| { builder .add_global_label("service", identity.service()) .add_global_label("version", identity.version()) .add_global_label("environment", identity.environment()) }) .map_err(|_| MetricsSurfaceError::RecorderConfiguration) } async fn render_metrics(State(state): State, headers: HeaderMap) -> Response { let Some(_permit) = ScrapePermit::try_acquire(&state.active_scrapes) else { return (StatusCode::SERVICE_UNAVAILABLE, TOO_MANY_SCRAPES).into_response(); }; let legacy = state.handle.render(); let openmetrics = headers .get(header::ACCEPT) .and_then(|value| value.to_str().ok()) .is_some_and(prefers_openmetrics); let body = if openmetrics { render_openmetrics(&legacy, &exemplar_snapshot()) } else { legacy }; exposition_response(body, openmetrics) } fn exposition_response(body: String, openmetrics: bool) -> Response { if body.len() > MAX_EXPOSITION_BYTES { return (StatusCode::SERVICE_UNAVAILABLE, EXPOSITION_TOO_LARGE).into_response(); } let mut response = body.into_response(); response.headers_mut().insert( header::CONTENT_TYPE, HeaderValue::from_static(if openmetrics { OPENMETRICS_CONTENT_TYPE } else { PROMETHEUS_CONTENT_TYPE }), ); response } fn render_openmetrics(legacy: &str, exemplars: &[Arc]) -> String { let index: HashMap = exemplars .iter() .map(|exemplar| (exemplar_key(exemplar), exemplar.as_ref())) .collect(); let mut output = String::with_capacity(legacy.len() + exemplars.len().saturating_mul(96) + 6); for line in legacy.lines() { output.push_str(line); if let Some(exemplar) = sample_key(line).and_then(|key| index.get(&key).copied()) { output.push_str(" # {trace_id=\""); output.push_str(exemplar.trace_id.as_str()); output.push_str("\"} "); output.push_str(&exemplar.value.to_string()); } output.push('\n'); } output.push_str("# EOF\n"); output } fn exemplar_key(exemplar: &ExemplarObservation) -> String { let mut labels = exemplar.labels.clone(); labels.sort_unstable(); let bound = exemplar .bucket_upper_bound .map_or_else(|| "+Inf".to_owned(), |value| value.to_string()); format!("{}|{:?}|{bound}", exemplar.metric, labels) } fn sample_key(line: &str) -> Option { let (head, _) = line.split_once("} ")?; let (metric, raw_labels) = head.split_once("_bucket{")?; let mut labels = Vec::new(); let mut bound = None; for item in raw_labels.split(',') { let (name, value) = item.split_once("=\"")?; let value = value.strip_suffix('"')?; match name { "le" => bound = Some(value), "service" | "version" | "environment" => {} _ => labels.push((name, value)), } } labels.sort_unstable(); Some(format!("{metric}|{labels:?}|{}", bound?)) } fn prefers_openmetrics(value: &str) -> bool { let mut open_q = 0.0_f32; let mut legacy_q = 0.0_f32; for range in value.split(',') { let mut parts = range.trim().split(';'); let media = parts.next().unwrap_or_default().trim(); let mut q = 1.0_f32; let mut supported_version = true; for parameter in parts { let Some((name, raw)) = parameter.trim().split_once('=') else { continue; }; match name.trim() { "q" => q = raw.trim().parse().unwrap_or(0.0), "version" if media == "application/openmetrics-text" => { supported_version = raw.trim().trim_matches('"') == "1.0.0"; } _ => {} } } if media == "application/openmetrics-text" && supported_version { open_q = open_q.max(q.clamp(0.0, 1.0)); } else if media == "text/plain" || media == "*/*" { legacy_q = legacy_q.max(q.clamp(0.0, 1.0)); } } open_q > 0.0 && open_q >= legacy_q } struct ScrapePermit(Arc); impl ScrapePermit { fn try_acquire(active: &Arc) -> Option { active .fetch_update(Ordering::AcqRel, Ordering::Acquire, |count| { (count < MAX_CONCURRENT_SCRAPES).then_some(count + 1) }) .ok() .map(|_| Self(Arc::clone(active))) } } impl Drop for ScrapePermit { fn drop(&mut self) { self.0.fetch_sub(1, Ordering::Release); } } async fn metrics_health() -> impl IntoResponse { (StatusCode::OK, "ok\n") } async fn authorize_metrics( State(state): State, request: Request, next: Next, ) -> Response { if !state.requires_authentication { return next.run(request).await; } let authorized = bearer_token(request.headers()) .map(token_digest) .zip(state.token_digest) .is_some_and(|(actual, expected)| bool::from(actual.ct_eq(&expected))); if authorized { next.run(request).await } else { StatusCode::UNAUTHORIZED.into_response() } } fn bearer_token(headers: &HeaderMap) -> Option<&[u8]> { let value = headers.get(header::AUTHORIZATION)?.as_bytes(); let separator = value.iter().position(|byte| *byte == b' ')?; let (scheme, token_with_spaces) = value.split_at(separator); let token_start = token_with_spaces.iter().position(|byte| *byte != b' ')?; let token = &token_with_spaces[token_start..]; scheme .eq_ignore_ascii_case(b"bearer") .then_some(token) .filter(|token| !token.is_empty()) } fn token_digest(token: &[u8]) -> [u8; 32] { Sha256::digest(token).into() } #[cfg(test)] mod rendering_tests { use std::{ sync::{Arc, atomic::AtomicUsize}, time::{Duration, Instant}, }; use crank_metrics::{ExemplarObservation, ExemplarTraceId}; use super::{ MAX_CONCURRENT_SCRAPES, ScrapePermit, exposition_response, prefers_openmetrics, render_openmetrics, }; #[test] fn openmetrics_adds_bounded_exemplar_without_changing_aggregate() { let legacy = "# TYPE crank_http_request_duration_seconds histogram\ncrank_http_request_duration_seconds_bucket{method=\"GET\",route=\"/health\",le=\"0.01\"} 1\ncrank_http_request_duration_seconds_sum{method=\"GET\",route=\"/health\"} 0.007\n"; let exemplar = ExemplarObservation { metric: "crank_http_request_duration_seconds", labels: vec![("route", "/health"), ("method", "GET")], bucket_upper_bound: Some(0.01), value: 0.007, trace_id: ExemplarTraceId::parse("0123456789abcdef0123456789abcdef").unwrap(), }; let rendered = render_openmetrics(legacy, &[Arc::new(exemplar)]); assert!(rendered.contains("# {trace_id=\"0123456789abcdef0123456789abcdef\"} 0.007")); assert!(rendered.ends_with("# EOF\n")); assert_eq!(rendered.matches(" 1").count(), legacy.matches(" 1").count()); } #[test] fn accept_negotiation_is_exact_and_honors_quality_and_version() { assert!(prefers_openmetrics( "application/openmetrics-text; version=1.0.0" )); assert!(prefers_openmetrics("application/openmetrics-text")); assert!(!prefers_openmetrics("application/openmetrics-text; q=0")); assert!(!prefers_openmetrics("application/openmetrics-textual")); assert!(!prefers_openmetrics( "application/openmetrics-text; version=0.0.1" )); assert!(!prefers_openmetrics( "application/openmetrics-text;q=0.2,text/plain;q=0.8" )); } #[test] fn output_and_concurrent_scrapes_fail_closed_at_their_bounds() { assert_eq!( exposition_response("x".repeat(crank_metrics::MAX_EXPOSITION_BYTES + 1), false) .status(), axum::http::StatusCode::SERVICE_UNAVAILABLE ); let active = Arc::new(AtomicUsize::new(0)); let permits = (0..MAX_CONCURRENT_SCRAPES) .map(|_| ScrapePermit::try_acquire(&active).unwrap()) .collect::>(); assert!(ScrapePermit::try_acquire(&active).is_none()); drop(permits); assert!(ScrapePermit::try_acquire(&active).is_some()); } #[test] fn maximum_fixture_renders_in_linear_bounded_time() { let trace = ExemplarTraceId::parse("0123456789abcdef0123456789abcdef").unwrap(); let mut legacy = String::new(); let mut exemplars = Vec::new(); for index in 0..5_000 { let value: &'static str = Box::leak(format!("route-{index}").into_boxed_str()); legacy.push_str(&format!( "crank_http_request_duration_seconds_bucket{{route=\"{value}\",method=\"GET\",le=\"0.01\"}} 1\n" )); exemplars.push(Arc::new(ExemplarObservation { metric: "crank_http_request_duration_seconds", labels: vec![("route", value), ("method", "GET")], bucket_upper_bound: Some(0.01), value: 0.007, trace_id: trace, })); } let started = Instant::now(); let rendered = render_openmetrics(&legacy, &exemplars); assert_eq!(rendered.matches("# {trace_id=").count(), 5_000); assert!(started.elapsed() < Duration::from_secs(3)); assert!(rendered.len() < crank_metrics::MAX_EXPOSITION_BYTES); } }