use std::{
borrow::Cow,
collections::{BTreeMap, BTreeSet},
fmt::Write as _,
future::Future,
pin::Pin,
sync::{Arc, Mutex},
task::{Context, Poll},
time::{Duration, Instant},
};
use axum::{Router, routing::get};
use http::{Method, Request, Response, StatusCode};
use tower::{Layer, Service};
pub fn metrics_layer<H>(hook: H) -> MetricsLayer<H>
where
H: HttpMetricsHook,
{
MetricsLayer::new(hook)
}
pub fn route_metrics_layer<H>(route: impl Into<Cow<'static, str>>, hook: H) -> MetricsLayer<H>
where
H: HttpMetricsHook,
{
MetricsLayer::new(hook).route(route)
}
pub trait HttpMetricsHook: Clone + Send + Sync + 'static {
fn on_request(&self, method: &Method, route: Option<&str>);
fn on_response(
&self,
method: &Method,
route: Option<&str>,
status: StatusCode,
latency: Duration,
);
fn on_error(&self, _method: &Method, _route: Option<&str>, _latency: Duration) {}
}
#[derive(Clone, Debug)]
pub struct PrometheusMetrics {
state: Arc<Mutex<PrometheusState>>,
excluded_routes: Arc<BTreeSet<String>>,
max_series: Option<usize>,
}
impl PrometheusMetrics {
pub fn new() -> Self {
Self {
state: Arc::new(Mutex::new(PrometheusState::default())),
excluded_routes: Arc::new(BTreeSet::from([
"/health/live".to_owned(),
"/health/ready".to_owned(),
"/metrics".to_owned(),
])),
max_series: None,
}
}
pub fn exclude_route(mut self, route: impl Into<String>) -> Self {
Arc::make_mut(&mut self.excluded_routes).insert(route.into());
self
}
pub fn with_max_series(mut self, max_series: usize) -> Self {
self.max_series = Some(max_series);
self
}
pub fn layer(&self) -> MetricsLayer<Self> {
MetricsLayer::new(self.clone())
}
pub fn routes(&self) -> Router {
self.routes_at("/metrics")
}
pub fn routes_at(&self, path: &'static str) -> Router {
let metrics = self.clone();
Router::new().route(path, get(move || async move { metrics.render() }))
}
pub fn render(&self) -> String {
let state = self.snapshot();
render_prometheus(&state)
}
fn snapshot(&self) -> PrometheusState {
self.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.clone()
}
fn should_record(&self, route: Option<&str>) -> bool {
route
.map(|route| !self.excluded_routes.contains(route))
.unwrap_or(true)
}
}
fn render_prometheus(state: &PrometheusState) -> String {
let mut output = String::new();
output.push_str("# TYPE nidus_http_requests_total counter\n");
for ((method, route, status), series) in &state.series {
let _ = writeln!(
output,
"nidus_http_requests_total{{method=\"{}\",route=\"{}\",status=\"{}\"}} {}",
escape_label(method),
escape_label(route),
status,
series.requests
);
}
output.push_str("# TYPE nidus_http_request_duration_seconds histogram\n");
for ((method, route, status), series) in &state.series {
let histogram = &series.histogram;
let method = escape_label(method);
let route = escape_label(route);
for (bucket, count) in HTTP_DURATION_BUCKET_LABELS
.iter()
.zip(histogram.bucket_counts.iter())
{
let _ = writeln!(
output,
"nidus_http_request_duration_seconds_bucket{{method=\"{method}\",route=\"{route}\",status=\"{status}\",le=\"{bucket}\"}} {count}",
);
}
let _ = writeln!(
output,
"nidus_http_request_duration_seconds_bucket{{method=\"{method}\",route=\"{route}\",status=\"{status}\",le=\"+Inf\"}} {}",
histogram.count
);
let _ = writeln!(
output,
"nidus_http_request_duration_seconds_count{{method=\"{method}\",route=\"{route}\",status=\"{status}\"}} {}",
histogram.count
);
let _ = writeln!(
output,
"nidus_http_request_duration_seconds_sum{{method=\"{method}\",route=\"{route}\",status=\"{status}\"}} {:.6}",
histogram.sum
);
}
output.push_str("# TYPE nidus_http_in_flight_requests gauge\n");
for ((method, route), count) in &state.in_flight {
let _ = writeln!(
output,
"nidus_http_in_flight_requests{{method=\"{}\",route=\"{}\"}} {}",
escape_label(method),
escape_label(route),
count
);
}
output.push_str("# TYPE nidus_http_errors_total counter\n");
for ((method, route, status), series) in &state.series {
if series.errors == 0 {
continue;
}
let _ = writeln!(
output,
"nidus_http_errors_total{{method=\"{}\",route=\"{}\",status=\"{}\"}} {}",
escape_label(method),
escape_label(route),
status,
series.errors
);
}
output
}
impl Default for PrometheusMetrics {
fn default() -> Self {
Self::new()
}
}
impl HttpMetricsHook for PrometheusMetrics {
fn on_request(&self, method: &Method, route: Option<&str>) {
if !self.should_record(route) {
return;
}
let route = route.unwrap_or("<unknown>").to_owned();
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let route = match self.max_series {
Some(max) => state.admit_route(route, max),
None => route,
};
*state
.in_flight
.entry((method.as_str().to_owned(), route))
.or_default() += 1;
}
fn on_response(
&self,
method: &Method,
route: Option<&str>,
status: StatusCode,
latency: Duration,
) {
let is_error = status.is_client_error() || status.is_server_error();
self.record_completion(method, route, status.as_u16(), latency, is_error);
}
fn on_error(&self, method: &Method, route: Option<&str>, latency: Duration) {
self.record_completion(
method,
route,
StatusCode::INTERNAL_SERVER_ERROR.as_u16(),
latency,
true,
);
}
}
impl PrometheusMetrics {
fn record_completion(
&self,
method: &Method,
route: Option<&str>,
status: u16,
latency: Duration,
is_error: bool,
) {
if !self.should_record(route) {
return;
}
let method = method.as_str().to_owned();
let route = route.unwrap_or("<unknown>").to_owned();
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let route = match self.max_series {
Some(max) => state.admit_route(route, max),
None => route,
};
let key = (method, route);
if let Some(count) = state.in_flight.get_mut(&key) {
*count = count.saturating_sub(1);
}
let (method, route) = key;
let series = state.series.entry((method, route, status)).or_default();
series.requests += 1;
series.histogram.observe(latency);
if is_error {
series.errors += 1;
}
}
}
#[derive(Clone, Debug, Default)]
struct PrometheusState {
series: BTreeMap<(String, String, u16), StatusSeries>,
in_flight: BTreeMap<(String, String), u64>,
known_routes: BTreeSet<String>,
}
#[derive(Clone, Debug, Default)]
struct StatusSeries {
requests: u64,
errors: u64,
histogram: DurationHistogram,
}
impl PrometheusState {
fn admit_route(&mut self, route: String, max_series: usize) -> String {
if self.known_routes.contains(&route) {
route
} else if self.known_routes.len() < max_series {
self.known_routes.insert(route.clone());
route
} else {
"<overflow>".to_owned()
}
}
}
const HTTP_DURATION_BUCKETS: [f64; 11] = [
0.005, 0.010, 0.025, 0.050, 0.100, 0.250, 0.500, 1.000, 2.500, 5.000, 10.000,
];
const HTTP_DURATION_BUCKET_LABELS: [&str; 11] = [
"0.005", "0.01", "0.025", "0.05", "0.1", "0.25", "0.5", "1", "2.5", "5", "10",
];
#[derive(Clone, Debug, Default)]
struct DurationHistogram {
count: u64,
sum: f64,
bucket_counts: [u64; HTTP_DURATION_BUCKETS.len()],
}
impl DurationHistogram {
fn observe(&mut self, latency: Duration) {
let seconds = latency.as_secs_f64();
self.count += 1;
self.sum += seconds;
for (bucket, count) in HTTP_DURATION_BUCKETS
.iter()
.zip(self.bucket_counts.iter_mut())
{
if seconds <= *bucket {
*count += 1;
}
}
}
}
#[derive(Clone, Debug)]
pub struct MetricsLayer<H> {
hook: H,
route: Option<Cow<'static, str>>,
}
impl<H> MetricsLayer<H>
where
H: HttpMetricsHook,
{
pub fn new(hook: H) -> Self {
Self { hook, route: None }
}
pub fn route(mut self, route: impl Into<Cow<'static, str>>) -> Self {
self.route = Some(route.into());
self
}
}
impl<S, H> Layer<S> for MetricsLayer<H>
where
H: HttpMetricsHook,
{
type Service = MetricsService<S, H>;
fn layer(&self, inner: S) -> Self::Service {
MetricsService {
inner,
hook: self.hook.clone(),
route: self.route.clone(),
}
}
}
#[derive(Clone, Debug)]
pub struct MetricsService<S, H> {
inner: S,
hook: H,
route: Option<Cow<'static, str>>,
}
impl<S, H, RequestBody, ResponseBody> Service<Request<RequestBody>> for MetricsService<S, H>
where
S: Service<Request<RequestBody>, Response = Response<ResponseBody>> + Send + 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
H: HttpMetricsHook,
RequestBody: Send + 'static,
ResponseBody: Send + 'static,
{
type Response = Response<ResponseBody>;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request<RequestBody>) -> Self::Future {
let method = request.method().clone();
let hook = self.hook.clone();
let route = match self.route.clone() {
Some(route) => RouteLabel::Fixed(route),
None => match request.extensions().get::<axum::extract::MatchedPath>() {
Some(path) => RouteLabel::Matched(path.clone()),
None => RouteLabel::Unknown,
},
};
hook.on_request(&method, route.as_deref());
let started_at = Instant::now();
let future = self.inner.call(request);
Box::pin(async move {
match future.await {
Ok(response) => {
hook.on_response(
&method,
route.as_deref(),
response.status(),
started_at.elapsed(),
);
Ok(response)
}
Err(error) => {
hook.on_error(&method, route.as_deref(), started_at.elapsed());
Err(error)
}
}
})
}
}
enum RouteLabel {
Fixed(Cow<'static, str>),
Matched(axum::extract::MatchedPath),
Unknown,
}
impl RouteLabel {
fn as_deref(&self) -> Option<&str> {
match self {
Self::Fixed(route) => Some(route),
Self::Matched(path) => Some(path.as_str()),
Self::Unknown => None,
}
}
}
fn escape_label(value: &str) -> Cow<'_, str> {
if value.contains(['\\', '\n', '"']) {
Cow::Owned(
value
.replace('\\', r"\\")
.replace('\n', r"\n")
.replace('"', r#"\""#),
)
} else {
Cow::Borrowed(value)
}
}
#[cfg(test)]
fn format_bucket(bucket: f64) -> String {
if bucket.fract() == 0.0 {
format!("{bucket:.0}")
} else {
let formatted = format!("{bucket:.3}");
formatted.trim_end_matches('0').to_owned()
}
}
#[cfg(test)]
mod tests {
use super::{HTTP_DURATION_BUCKET_LABELS, HTTP_DURATION_BUCKETS, escape_label, format_bucket};
#[test]
fn duration_bucket_labels_match_formatted_buckets() {
for (bucket, label) in HTTP_DURATION_BUCKETS
.iter()
.zip(HTTP_DURATION_BUCKET_LABELS.iter())
{
assert_eq!(&format_bucket(*bucket), label);
}
}
#[test]
fn escape_label_borrows_clean_values_and_escapes_special_characters() {
assert!(matches!(
escape_label("/users/{id}"),
std::borrow::Cow::Borrowed("/users/{id}")
));
assert_eq!(escape_label("a\\b\nc\"d"), r#"a\\b\nc\"d"#);
}
}