rmcp_server_kit/
metrics.rs1use std::sync::Arc;
16
17use prometheus::{
18 Encoder, HistogramOpts, HistogramVec, IntCounterVec, Registry, TextEncoder, opts,
19};
20
21use crate::error::RmcpServerKitError;
22
23const HTTP_DURATION_BUCKETS: &[f64] = &[
30 0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0,
31];
32
33#[derive(Clone, Debug)]
35#[non_exhaustive]
36pub struct McpMetrics {
37 pub registry: Registry,
39 pub http_requests_total: IntCounterVec,
41 pub http_request_duration_seconds: HistogramVec,
43 pub rate_limited_total: IntCounterVec,
48}
49
50impl McpMetrics {
51 pub fn new() -> Result<Self, RmcpServerKitError> {
58 let registry = Registry::new();
59
60 let http_requests_total = IntCounterVec::new(
61 opts!("rmcp_server_kit_http_requests_total", "Total HTTP requests"),
62 &["method", "path", "status"],
63 )
64 .map_err(|e| RmcpServerKitError::Metrics(e.to_string()))?;
65 registry
66 .register(Box::new(http_requests_total.clone()))
67 .map_err(|e| RmcpServerKitError::Metrics(e.to_string()))?;
68
69 let http_request_duration_seconds = HistogramVec::new(
70 HistogramOpts::new(
71 "rmcp_server_kit_http_request_duration_seconds",
72 "HTTP request duration in seconds",
73 )
74 .buckets(HTTP_DURATION_BUCKETS.to_vec()),
75 &["method", "path"],
76 )
77 .map_err(|e| RmcpServerKitError::Metrics(e.to_string()))?;
78 registry
79 .register(Box::new(http_request_duration_seconds.clone()))
80 .map_err(|e| RmcpServerKitError::Metrics(e.to_string()))?;
81
82 let rate_limited_total = IntCounterVec::new(
83 opts!(
84 "rmcp_server_kit_rate_limited_total",
85 "Rate-limiter denials by limiter"
86 ),
87 &["limiter"],
88 )
89 .map_err(|e| RmcpServerKitError::Metrics(e.to_string()))?;
90 registry
91 .register(Box::new(rate_limited_total.clone()))
92 .map_err(|e| RmcpServerKitError::Metrics(e.to_string()))?;
93
94 Ok(Self {
95 registry,
96 http_requests_total,
97 http_request_duration_seconds,
98 rate_limited_total,
99 })
100 }
101
102 #[must_use]
104 pub fn encode(&self) -> String {
105 let encoder = TextEncoder::new();
106 let metric_families = self.registry.gather();
107 let mut buf = Vec::new();
108 if let Err(e) = encoder.encode(&metric_families, &mut buf) {
109 tracing::warn!(error = %e, "prometheus encode failed");
110 return String::new();
111 }
112 String::from_utf8(buf).unwrap_or_default()
115 }
116}
117
118pub(crate) fn record_rate_limit_deny(ext: &axum::http::Extensions, limiter: &str) {
127 if let Some(m) = ext.get::<Arc<McpMetrics>>() {
128 m.rate_limited_total.with_label_values(&[limiter]).inc();
129 }
130}
131
132pub async fn serve_metrics(
146 bind: String,
147 metrics: Arc<McpMetrics>,
148 shutdown: tokio_util::sync::CancellationToken,
149) -> Result<(), RmcpServerKitError> {
150 let app = axum::Router::new().route(
151 "/metrics",
152 axum::routing::get(move || {
153 let m = Arc::clone(&metrics);
154 async move { m.encode() }
155 }),
156 );
157
158 let listener = tokio::net::TcpListener::bind(&bind)
159 .await
160 .map_err(|e| RmcpServerKitError::Startup(format!("metrics bind {bind}: {e}")))?;
161 tracing::info!("metrics endpoint listening on http://{bind}/metrics");
162 axum::serve(listener, app)
163 .with_graceful_shutdown(async move { shutdown.cancelled().await })
164 .await
165 .map_err(|e| RmcpServerKitError::Startup(format!("metrics serve: {e}")))?;
166 Ok(())
167}
168
169#[cfg(test)]
170mod tests {
171 #![allow(
172 clippy::unwrap_used,
173 clippy::expect_used,
174 clippy::panic,
175 clippy::indexing_slicing,
176 clippy::unwrap_in_result,
177 clippy::print_stdout,
178 clippy::print_stderr,
179 reason = "test-only relaxations; production code uses ? and tracing"
180 )]
181 use super::*;
182
183 #[test]
184 fn new_creates_registry_with_counters() {
185 let m = McpMetrics::new().unwrap();
186 m.http_requests_total
188 .with_label_values(&["GET", "/test", "200"])
189 .inc();
190 m.http_request_duration_seconds
191 .with_label_values(&["GET", "/test"])
192 .observe(0.1);
193 assert_eq!(m.registry.gather().len(), 2);
194 }
195
196 #[test]
197 fn encode_empty_registry() {
198 let m = McpMetrics::new().unwrap();
199 let output = m.encode();
200 assert!(output.is_empty() || output.contains("rmcp_server_kit_"));
202 }
203
204 #[test]
205 fn counter_increment_shows_in_encode() {
206 let m = McpMetrics::new().unwrap();
207 m.http_requests_total
208 .with_label_values(&["GET", "/healthz", "200"])
209 .inc();
210 let output = m.encode();
211 assert!(output.contains("rmcp_server_kit_http_requests_total"));
212 assert!(output.contains("method=\"GET\""));
213 assert!(output.contains("path=\"/healthz\""));
214 assert!(output.contains("status=\"200\""));
215 assert!(output.contains(" 1")); }
217
218 #[test]
219 fn histogram_observe_shows_in_encode() {
220 let m = McpMetrics::new().unwrap();
221 m.http_request_duration_seconds
222 .with_label_values(&["POST", "/mcp"])
223 .observe(0.042);
224 let output = m.encode();
225 assert!(output.contains("rmcp_server_kit_http_request_duration_seconds"));
226 assert!(output.contains("method=\"POST\""));
227 assert!(output.contains("path=\"/mcp\""));
228 }
229
230 #[test]
231 fn multiple_increments_accumulate() {
232 let m = McpMetrics::new().unwrap();
233 let counter = m
234 .http_requests_total
235 .with_label_values(&["POST", "/mcp", "200"]);
236 counter.inc();
237 counter.inc();
238 counter.inc();
239 let output = m.encode();
240 assert!(output.contains(" 3")); }
242
243 #[test]
244 fn clone_shares_registry() {
245 let m = McpMetrics::new().unwrap();
246 let m2 = m.clone();
247 m.http_requests_total
248 .with_label_values(&["GET", "/test", "200"])
249 .inc();
250 let output = m2.encode();
252 assert!(output.contains(" 1"));
253 }
254
255 #[test]
256 fn rate_limited_counter_registers_and_encodes() {
257 let m = McpMetrics::new().unwrap();
258 m.rate_limited_total.with_label_values(&["tool"]).inc();
259 let output = m.encode();
260 assert!(output.contains("rmcp_server_kit_rate_limited_total"));
261 assert!(output.contains("limiter=\"tool\""));
262 assert!(output.contains(" 1"));
263 }
264
265 #[test]
266 fn record_rate_limit_deny_increments_via_extension() {
267 let m = Arc::new(McpMetrics::new().unwrap());
268 let mut ext = axum::http::Extensions::new();
269 ext.insert(Arc::clone(&m));
270 record_rate_limit_deny(&ext, "auth_pre");
271 record_rate_limit_deny(&ext, "auth_pre");
272 assert_eq!(
273 m.rate_limited_total.with_label_values(&["auth_pre"]).get(),
274 2
275 );
276 let empty = axum::http::Extensions::new();
278 record_rate_limit_deny(&empty, "auth_pre");
279 assert_eq!(
280 m.rate_limited_total.with_label_values(&["auth_pre"]).get(),
281 2
282 );
283 }
284
285 #[tokio::test]
291 async fn serve_metrics_releases_port_on_shutdown() {
292 let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
295 let addr = probe.local_addr().unwrap();
296 drop(probe);
297
298 let metrics = Arc::new(McpMetrics::new().unwrap());
299 let shutdown = tokio_util::sync::CancellationToken::new();
300 let handle = tokio::spawn(serve_metrics(
301 addr.to_string(),
302 Arc::clone(&metrics),
303 shutdown.clone(),
304 ));
305
306 let deadline = std::time::Instant::now() + std::time::Duration::from_secs(2);
308 loop {
309 if tokio::net::TcpStream::connect(addr).await.is_ok() {
310 break;
311 }
312 assert!(
313 std::time::Instant::now() < deadline,
314 "metrics listener never accepted on {addr}"
315 );
316 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
317 }
318
319 shutdown.cancel();
321 let join = tokio::time::timeout(std::time::Duration::from_secs(5), handle)
322 .await
323 .expect("serve_metrics did not return within timeout");
324 join.expect("join error")
325 .expect("serve_metrics returned Err");
326
327 let rebind = tokio::net::TcpListener::bind(addr)
329 .await
330 .expect("port not released after shutdown");
331 drop(rebind);
332 }
333}