use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant};
use crate::security::sanitize_text;
fn sanitize_metric_label(value: &str) -> String {
sanitize_text(value)
}
fn prometheus_label_value(value: &str) -> String {
sanitize_metric_label(value)
.replace('\\', r"\\")
.replace('\n', r"\n")
.replace('"', r#"\""#)
}
#[derive(Debug, Clone)]
pub struct MetricsConfig {
pub enabled: bool,
pub histogram_buckets: Vec<f64>,
pub endpoint_path: String,
}
impl Default for MetricsConfig {
fn default() -> Self {
Self {
enabled: true,
histogram_buckets: vec![
0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0,
],
endpoint_path: "/metrics".to_string(),
}
}
}
#[derive(Debug, Default)]
pub struct Counter {
value: AtomicU64,
}
impl Counter {
pub fn new() -> Self {
Self::default()
}
pub fn inc(&self) {
self.value.fetch_add(1, Ordering::Relaxed);
}
pub fn add(&self, v: u64) {
self.value.fetch_add(v, Ordering::Relaxed);
}
pub fn get(&self) -> u64 {
self.value.load(Ordering::Relaxed)
}
}
#[derive(Debug)]
pub struct Histogram {
buckets: Vec<f64>,
bucket_counts: Vec<AtomicU64>,
sum: AtomicU64,
count: AtomicU64,
}
impl Histogram {
pub fn new(buckets: Vec<f64>) -> Self {
let bucket_counts = buckets.iter().map(|_| AtomicU64::new(0)).collect();
Self {
buckets,
bucket_counts,
sum: AtomicU64::new(0),
count: AtomicU64::new(0),
}
}
pub fn with_default_buckets() -> Self {
Self::new(vec![
0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0,
])
}
pub fn observe(&self, value: f64) {
let micros = (value * 1_000_000.0) as u64;
self.sum.fetch_add(micros, Ordering::Relaxed);
self.count.fetch_add(1, Ordering::Relaxed);
for (i, &bound) in self.buckets.iter().enumerate() {
if value <= bound {
self.bucket_counts[i].fetch_add(1, Ordering::Relaxed);
}
}
}
pub fn observe_duration(&self, duration: Duration) {
self.observe(duration.as_secs_f64());
}
pub fn bucket_counts(&self) -> Vec<u64> {
self.bucket_counts
.iter()
.map(|c| c.load(Ordering::Relaxed))
.collect()
}
pub fn sum(&self) -> f64 {
self.sum.load(Ordering::Relaxed) as f64 / 1_000_000.0
}
pub fn count(&self) -> u64 {
self.count.load(Ordering::Relaxed)
}
pub fn buckets(&self) -> &[f64] {
&self.buckets
}
}
#[derive(Debug, Default)]
pub struct LabeledCounter {
counters: RwLock<HashMap<String, Arc<Counter>>>,
}
impl LabeledCounter {
pub fn new() -> Self {
Self::default()
}
pub fn with_label(&self, label: &str) -> Arc<Counter> {
{
let counters = self.counters.read().unwrap();
if let Some(counter) = counters.get(label) {
return Arc::clone(counter);
}
}
let mut counters = self.counters.write().unwrap();
counters
.entry(label.to_string())
.or_insert_with(|| Arc::new(Counter::new()))
.clone()
}
pub fn inc(&self, label: &str) {
self.with_label(label).inc();
}
pub fn all(&self) -> HashMap<String, u64> {
let counters = self.counters.read().unwrap();
counters.iter().map(|(k, v)| (k.clone(), v.get())).collect()
}
}
#[derive(Debug)]
pub struct LabeledHistogram {
histograms: RwLock<HashMap<String, Arc<Histogram>>>,
buckets: Vec<f64>,
}
impl LabeledHistogram {
pub fn new(buckets: Vec<f64>) -> Self {
Self {
histograms: RwLock::new(HashMap::new()),
buckets,
}
}
pub fn with_default_buckets() -> Self {
Self::new(vec![
0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0,
])
}
pub fn with_label(&self, label: &str) -> Arc<Histogram> {
{
let histograms = self.histograms.read().unwrap();
if let Some(histogram) = histograms.get(label) {
return Arc::clone(histogram);
}
}
let mut histograms = self.histograms.write().unwrap();
histograms
.entry(label.to_string())
.or_insert_with(|| Arc::new(Histogram::new(self.buckets.clone())))
.clone()
}
pub fn observe(&self, label: &str, value: f64) {
self.with_label(label).observe(value);
}
pub fn labels(&self) -> Vec<String> {
let histograms = self.histograms.read().unwrap();
histograms.keys().cloned().collect()
}
}
pub struct PrometheusMetrics {
pub requests_total: LabeledCounter,
pub request_duration_seconds: LabeledHistogram,
pub active_connections: Counter,
pub bytes_total: LabeledCounter,
pub errors_total: LabeledCounter,
config: MetricsConfig,
}
impl PrometheusMetrics {
pub fn new(config: MetricsConfig) -> Self {
Self {
requests_total: LabeledCounter::new(),
request_duration_seconds: LabeledHistogram::new(config.histogram_buckets.clone()),
active_connections: Counter::new(),
bytes_total: LabeledCounter::new(),
errors_total: LabeledCounter::new(),
config,
}
}
pub fn with_defaults() -> Self {
Self::new(MetricsConfig::default())
}
pub fn record_request(&self, method: &str, duration: Duration) {
let method = sanitize_metric_label(method);
self.requests_total.inc(&method);
self.request_duration_seconds
.observe(&method, duration.as_secs_f64());
}
pub fn record_error(&self, error_type: &str) {
self.errors_total.inc(&sanitize_metric_label(error_type));
}
pub fn record_bytes(&self, direction: &str, bytes: u64) {
self.bytes_total
.with_label(&sanitize_metric_label(direction))
.add(bytes);
}
pub fn connection_opened(&self) {
self.active_connections.inc();
}
pub fn connection_closed(&self) {
}
pub fn config(&self) -> &MetricsConfig {
&self.config
}
pub fn format_prometheus(&self) -> String {
let mut output = String::new();
output.push_str("# HELP dcp_requests_total Total number of requests by method\n");
output.push_str("# TYPE dcp_requests_total counter\n");
for (method, count) in self.requests_total.all() {
let method = prometheus_label_value(&method);
output.push_str(&format!(
"dcp_requests_total{{method=\"{}\"}} {}\n",
method, count
));
}
output.push_str("\n# HELP dcp_request_duration_seconds Request latency in seconds\n");
output.push_str("# TYPE dcp_request_duration_seconds histogram\n");
for label in self.request_duration_seconds.labels() {
let histogram = self.request_duration_seconds.with_label(&label);
let buckets = histogram.buckets();
let counts = histogram.bucket_counts();
let method = prometheus_label_value(&label);
for (i, &bound) in buckets.iter().enumerate() {
output.push_str(&format!(
"dcp_request_duration_seconds_bucket{{method=\"{}\",le=\"{}\"}} {}\n",
method, bound, counts[i]
));
}
output.push_str(&format!(
"dcp_request_duration_seconds_bucket{{method=\"{}\",le=\"+Inf\"}} {}\n",
method,
histogram.count()
));
output.push_str(&format!(
"dcp_request_duration_seconds_sum{{method=\"{}\"}} {}\n",
method,
histogram.sum()
));
output.push_str(&format!(
"dcp_request_duration_seconds_count{{method=\"{}\"}} {}\n",
method,
histogram.count()
));
}
output.push_str("\n# HELP dcp_active_connections Number of active connections\n");
output.push_str("# TYPE dcp_active_connections gauge\n");
output.push_str(&format!(
"dcp_active_connections {}\n",
self.active_connections.get()
));
output.push_str("\n# HELP dcp_bytes_total Total bytes transferred\n");
output.push_str("# TYPE dcp_bytes_total counter\n");
for (direction, count) in self.bytes_total.all() {
let direction = prometheus_label_value(&direction);
output.push_str(&format!(
"dcp_bytes_total{{direction=\"{}\"}} {}\n",
direction, count
));
}
output.push_str("\n# HELP dcp_errors_total Total number of errors by type\n");
output.push_str("# TYPE dcp_errors_total counter\n");
for (error_type, count) in self.errors_total.all() {
let error_type = prometheus_label_value(&error_type);
output.push_str(&format!(
"dcp_errors_total{{type=\"{}\"}} {}\n",
error_type, count
));
}
output
}
}
impl Default for PrometheusMetrics {
fn default() -> Self {
Self::with_defaults()
}
}
pub struct RequestMetrics {
method: String,
start: Instant,
metrics: Arc<PrometheusMetrics>,
}
impl RequestMetrics {
pub fn start(metrics: Arc<PrometheusMetrics>, method: impl Into<String>) -> Self {
Self {
method: method.into(),
start: Instant::now(),
metrics,
}
}
pub fn finish(self) {
let duration = self.start.elapsed();
self.metrics.record_request(&self.method, duration);
}
pub fn finish_with_error(self, error_type: &str) {
let duration = self.start.elapsed();
self.metrics.record_request(&self.method, duration);
self.metrics.record_error(error_type);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_counter() {
let counter = Counter::new();
assert_eq!(counter.get(), 0);
counter.inc();
assert_eq!(counter.get(), 1);
counter.add(5);
assert_eq!(counter.get(), 6);
}
#[test]
fn test_histogram() {
let histogram = Histogram::new(vec![0.1, 0.5, 1.0]);
histogram.observe(0.05);
histogram.observe(0.3);
histogram.observe(0.8);
histogram.observe(2.0);
assert_eq!(histogram.count(), 4);
let counts = histogram.bucket_counts();
assert_eq!(counts[0], 1); assert_eq!(counts[1], 2); assert_eq!(counts[2], 3); }
#[test]
fn test_labeled_counter() {
let counter = LabeledCounter::new();
counter.inc("method_a");
counter.inc("method_a");
counter.inc("method_b");
let all = counter.all();
assert_eq!(all.get("method_a"), Some(&2));
assert_eq!(all.get("method_b"), Some(&1));
}
#[test]
fn test_prometheus_metrics() {
let metrics = PrometheusMetrics::with_defaults();
metrics.record_request("tools/list", Duration::from_millis(10));
metrics.record_request("tools/call", Duration::from_millis(50));
metrics.record_error("timeout");
metrics.record_bytes("in", 1000);
metrics.record_bytes("out", 500);
let output = metrics.format_prometheus();
assert!(output.contains("dcp_requests_total"));
assert!(output.contains("dcp_errors_total"));
assert!(output.contains("dcp_bytes_total"));
}
#[test]
fn test_request_metrics() {
let metrics = Arc::new(PrometheusMetrics::with_defaults());
let request = RequestMetrics::start(Arc::clone(&metrics), "test_method");
std::thread::sleep(Duration::from_millis(1));
request.finish();
let all = metrics.requests_total.all();
assert_eq!(all.get("test_method"), Some(&1));
}
}