Support in-flight request age metrics for router (#16341)

This commit is contained in:
fzyzcjy
2026-01-04 07:38:53 +08:00
committed by GitHub
parent e139d2aa76
commit c88aaf22c7
8 changed files with 428 additions and 9 deletions
+3
View File
@@ -14,6 +14,7 @@ use crate::{
},
mcp::McpManager,
middleware::TokenBucket,
observability::inflight_tracker::InFlightRequestTracker,
policies::PolicyRegistry,
reasoning_parser::ParserFactory as ReasoningParserFactory,
routers::router_manager::RouterManager,
@@ -62,6 +63,7 @@ pub struct AppContext {
pub mcp_manager: Arc<OnceLock<Arc<McpManager>>>,
pub wasm_manager: Option<Arc<WasmModuleManager>>,
pub worker_service: Arc<WorkerService>,
pub inflight_tracker: Arc<InFlightRequestTracker>,
}
pub struct AppContextBuilder {
@@ -275,6 +277,7 @@ impl AppContextBuilder {
.ok_or(AppContextBuildError("mcp_manager"))?,
wasm_manager: self.wasm_manager,
worker_service,
inflight_tracker: InFlightRequestTracker::new(),
})
}
+20 -6
View File
@@ -26,7 +26,10 @@ use tracing::{debug, error, field::Empty, info, info_span, warn, Span};
pub use crate::core::token_bucket::TokenBucket;
use crate::{
observability::metrics::{method_to_static_str, metrics_labels, Metrics},
observability::{
inflight_tracker::InFlightRequestTracker,
metrics::{method_to_static_str, metrics_labels, Metrics},
},
routers::error::extract_error_code_from_response,
server::AppState,
wasm::{
@@ -607,12 +610,14 @@ pub async fn concurrency_limit_middleware(
static ACTIVE_HTTP_CONNECTIONS: AtomicU64 = AtomicU64::new(0);
/// Tower Layer for HTTP metrics collection (SMG Layer 1 metrics)
#[derive(Clone, Copy, Default)]
pub struct HttpMetricsLayer;
#[derive(Clone)]
pub struct HttpMetricsLayer {
tracker: Arc<InFlightRequestTracker>,
}
impl HttpMetricsLayer {
pub fn new() -> Self {
Self
pub fn new(tracker: Arc<InFlightRequestTracker>) -> Self {
Self { tracker }
}
}
@@ -620,7 +625,10 @@ impl<S> Layer<S> for HttpMetricsLayer {
type Service = HttpMetricsMiddleware<S>;
fn layer(&self, inner: S) -> Self::Service {
HttpMetricsMiddleware { inner }
HttpMetricsMiddleware {
inner,
in_flight_request_tracker: self.tracker.clone(),
}
}
}
@@ -628,6 +636,7 @@ impl<S> Layer<S> for HttpMetricsLayer {
#[derive(Clone)]
pub struct HttpMetricsMiddleware<S> {
inner: S,
in_flight_request_tracker: Arc<InFlightRequestTracker>,
}
impl<S> Service<Request> for HttpMetricsMiddleware<S>
@@ -651,15 +660,20 @@ where
let start = Instant::now();
let mut inner = self.inner.clone();
let in_flight_request_tracker = self.in_flight_request_tracker.clone();
Box::pin(async move {
// Increment inside async block - ensures no leak if future is dropped before polling
let active = ACTIVE_HTTP_CONNECTIONS.fetch_add(1, Ordering::Relaxed) + 1;
Metrics::set_http_connections_active(active as usize);
let guard = in_flight_request_tracker.track();
// Capture result before decrementing to ensure decrement happens on error too
let result = inner.call(req).await;
drop(guard);
// Always decrement, regardless of success or failure
let active = ACTIVE_HTTP_CONNECTIONS.fetch_sub(1, Ordering::Relaxed) - 1;
Metrics::set_http_connections_active(active as usize);
@@ -0,0 +1,225 @@
use std::{
sync::{
atomic::{AtomicU64, Ordering},
Arc, OnceLock,
},
time::Instant,
};
use dashmap::DashMap;
use super::metrics::Metrics;
use crate::policies::utils::PeriodicTask;
const AGE_BUCKET_BOUNDS: &[u64] = &[30, 60, 180, 300, 600, 1200, 3600, 7200, 14400, 28800, 86400];
const AGE_BUCKET_LABELS: &[&str] = &[
"30", "60", "180", "300", "600", "1200", "3600", "7200", "14400", "28800", "86400", "+Inf",
];
pub struct InFlightRequestTracker {
requests: DashMap<u64, Instant>,
next_id: AtomicU64,
sampler: OnceLock<PeriodicTask>,
}
impl InFlightRequestTracker {
pub fn new() -> Arc<Self> {
Arc::new(Self {
requests: DashMap::new(),
next_id: AtomicU64::new(0),
sampler: OnceLock::new(),
})
}
pub fn start_sampler(self: &Arc<Self>, interval_secs: u64) {
let tracker = self.clone();
let task = PeriodicTask::spawn(interval_secs, "InFlightRequestSampler", move || {
tracker.sample_and_record();
});
self.sampler.set(task).unwrap();
}
pub fn track(self: &Arc<Self>) -> InFlightGuard {
let request_id = self.next_id.fetch_add(1, Ordering::Relaxed);
self.requests.insert(request_id, Instant::now());
InFlightGuard {
tracker: self.clone(),
request_id,
}
}
pub fn len(&self) -> usize {
self.requests.len()
}
pub fn is_empty(&self) -> bool {
self.requests.is_empty()
}
pub fn compute_bucket_counts(&self) -> [usize; AGE_BUCKET_LABELS.len()] {
let now = Instant::now();
let inf_idx = AGE_BUCKET_LABELS.len() - 1;
let instants: Vec<Instant> = self.requests.iter().map(|entry| *entry.value()).collect();
let mut non_cumulative_counts = [0usize; AGE_BUCKET_LABELS.len()];
for inst in instants {
let age_secs = now.duration_since(inst).as_secs();
let bucket_idx = AGE_BUCKET_BOUNDS
.iter()
.position(|&bound| age_secs <= bound)
.unwrap_or(inf_idx);
non_cumulative_counts[bucket_idx] += 1;
}
let mut counts = [0usize; AGE_BUCKET_LABELS.len()];
let mut cumulative = 0;
for i in 0..counts.len() {
cumulative += non_cumulative_counts[i];
counts[i] = cumulative;
}
counts
}
fn sample_and_record(&self) {
let counts = self.compute_bucket_counts();
for (i, &label) in AGE_BUCKET_LABELS.iter().enumerate() {
Metrics::set_inflight_request_age_count(label, counts[i]);
}
}
}
pub struct InFlightGuard {
tracker: Arc<InFlightRequestTracker>,
request_id: u64,
}
impl Drop for InFlightGuard {
fn drop(&mut self) {
self.tracker.requests.remove(&self.request_id);
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
impl InFlightRequestTracker {
fn insert_with_time(&self, request_id: u64, start_time: Instant) {
self.requests.insert(request_id, start_time);
}
}
#[test]
fn test_track_and_drop() {
let tracker = InFlightRequestTracker::new();
{
let _guard1 = tracker.track();
let _guard2 = tracker.track();
assert_eq!(tracker.len(), 2);
}
assert_eq!(tracker.len(), 0);
}
#[test]
fn test_guard_auto_deregister() {
let tracker = InFlightRequestTracker::new();
let guard = tracker.track();
assert_eq!(tracker.len(), 1);
drop(guard);
assert_eq!(tracker.len(), 0);
}
#[test]
fn test_request_age_tracking() {
let tracker = InFlightRequestTracker::new();
let _guard = tracker.track();
std::thread::sleep(Duration::from_millis(100));
let entry = tracker.requests.iter().next().unwrap();
let age = entry.value().elapsed();
assert!(age >= Duration::from_millis(100));
}
#[test]
fn test_empty_tracker_buckets() {
let tracker = InFlightRequestTracker::new();
let counts = tracker.compute_bucket_counts();
assert!(counts.iter().all(|&c| c == 0));
}
#[test]
fn test_cumulative_bucket_counts() {
let tracker = InFlightRequestTracker::new();
let now = Instant::now();
tracker.insert_with_time(1, now);
tracker.insert_with_time(2, now - Duration::from_secs(45));
tracker.insert_with_time(3, now - Duration::from_secs(100));
tracker.insert_with_time(4, now - Duration::from_secs(250));
tracker.insert_with_time(5, now - Duration::from_secs(500));
tracker.insert_with_time(6, now - Duration::from_secs(700));
let counts = tracker.compute_bucket_counts();
assert_eq!(counts[0], 1, "bucket 0");
assert_eq!(counts[1], 2, "bucket 1");
assert_eq!(counts[2], 3, "bucket 2");
assert_eq!(counts[3], 4, "bucket 3");
assert_eq!(counts[4], 5, "bucket 4");
assert_eq!(counts[5], 6, "bucket 5");
assert_eq!(*counts.last().unwrap(), 6, "bucket +Inf");
}
#[test]
fn test_bucket_boundary_values() {
let tracker = InFlightRequestTracker::new();
let now = Instant::now();
tracker.insert_with_time(1, now - Duration::from_secs(30));
tracker.insert_with_time(2, now - Duration::from_secs(31));
let counts = tracker.compute_bucket_counts();
assert_eq!(counts[0], 1, "bucket 0 includes exact boundary");
assert_eq!(counts[1], 2, "bucket 1 includes both");
assert_eq!(*counts.last().unwrap(), 2, "bucket +Inf includes all");
}
#[test]
fn test_concurrent_tracking() {
use std::thread;
let tracker = InFlightRequestTracker::new();
let mut handles = vec![];
for _ in 0..10 {
let t = tracker.clone();
handles.push(thread::spawn(move || {
(0..100).map(|_| t.track()).collect::<Vec<_>>()
}));
}
let all_guards: Vec<_> = handles
.into_iter()
.flat_map(|h| h.join().unwrap())
.collect();
assert_eq!(tracker.len(), 1000);
drop(all_guards);
assert_eq!(tracker.len(), 0);
}
#[test]
fn test_unique_ids() {
let tracker = InFlightRequestTracker::new();
let g1 = tracker.track();
let g2 = tracker.track();
let g3 = tracker.track();
assert_ne!(g1.request_id, g2.request_id);
assert_ne!(g2.request_id, g3.request_id);
assert_eq!(tracker.len(), 3);
}
}
+15 -1
View File
@@ -153,6 +153,10 @@ pub fn init_metrics() {
"smg_http_request_duration_seconds",
"HTTP request duration by method and path"
);
describe_gauge!(
"smg_http_inflight_request_age_count",
"Count of currently in-flight HTTP requests by age"
);
describe_counter!(
"smg_http_responses_total",
"Total HTTP responses by status_code and error_code"
@@ -491,7 +495,17 @@ impl Metrics {
.record(duration.as_secs_f64());
}
/// Set active HTTP connections count.
/// Set the cumulative count of in-flight requests for a given age bucket.
/// Uses `le` label to match Prometheus histogram convention.
pub fn set_inflight_request_age_count(le: &'static str, count: usize) {
gauge!(
"smg_http_inflight_request_age_count",
"le" => le
)
.set(count as f64);
}
/// Set active HTTP connections count
pub fn set_http_connections_active(count: usize) {
gauge!("smg_http_connections_active").set(count as f64);
}
@@ -1,6 +1,7 @@
//! Observability utilities for logging, metrics, and tracing.
pub mod events;
pub mod inflight_tracker;
pub mod logging;
pub mod metrics;
pub mod otel_trace;
+7 -1
View File
@@ -654,7 +654,9 @@ pub fn build_app(
max_payload_size,
))
.layer(middleware::create_logging_layer())
.layer(middleware::HttpMetricsLayer::new())
.layer(middleware::HttpMetricsLayer::new(
app_state.context.inflight_tracker.clone(),
))
.layer(middleware::RequestIdLayer::new(request_id_headers))
.layer(create_cors_layer(cors_allowed_origins))
.fallback(sink_handler)
@@ -714,6 +716,10 @@ pub async fn startup(config: ServerConfig) -> Result<(), Box<dyn std::error::Err
AppContext::from_config(config.router_config.clone(), config.request_timeout_secs).await?,
);
if config.prometheus_config.is_some() {
app_context.inflight_tracker.start_sampler(20);
}
let weak_context = Arc::downgrade(&app_context);
let worker_job_queue = JobQueue::new(JobQueueConfig::default(), weak_context);
app_context
+5 -1
View File
@@ -609,7 +609,10 @@ mod tests {
}
async fn create_test_app_context() -> Arc<AppContext> {
use crate::{config::RouterConfig, core::WorkerService, middleware::TokenBucket};
use crate::{
config::RouterConfig, core::WorkerService, middleware::TokenBucket,
observability::inflight_tracker::InFlightRequestTracker,
};
let router_config = RouterConfig::builder()
.worker_startup_timeout_secs(1)
@@ -649,6 +652,7 @@ mod tests {
worker_job_queue,
router_config,
)),
inflight_tracker: InFlightRequestTracker::new(),
})
}
@@ -0,0 +1,152 @@
mod common;
use std::time::Duration;
use axum::{
body::Body,
extract::Request,
http::{header::CONTENT_TYPE, StatusCode},
};
use common::{
mock_worker::{HealthStatus, MockWorkerConfig, WorkerType},
AppTestContext,
};
use serde_json::json;
use tower::ServiceExt;
#[tokio::test]
async fn test_multiple_concurrent_requests_tracking() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 19002,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 50,
fail_rate: 0.0,
}])
.await;
let tracker = ctx.app_context.inflight_tracker.clone();
let mut handles = vec![];
for i in 0..5 {
let app = ctx.create_app().await;
handles.push(tokio::spawn(async move {
let payload = json!({
"text": format!("Request {}", i),
"stream": false
});
let req = Request::builder()
.method("POST")
.uri("/generate")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
app.oneshot(req).await.unwrap()
}));
}
for handle in handles {
let resp = handle.await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
assert!(tracker.is_empty());
ctx.shutdown().await;
}
#[tokio::test]
async fn test_inflight_request_appears_in_bucket() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 19004,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 2000,
fail_rate: 0.0,
}])
.await;
let tracker = &ctx.app_context.inflight_tracker;
assert!(tracker.is_empty(), "Tracker should start empty");
let app = ctx.create_app().await;
let payload = json!({
"text": "Long running request",
"stream": false
});
let req = Request::builder()
.method("POST")
.uri("/generate")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let tracker_clone = ctx.app_context.inflight_tracker.clone();
let response_future = tokio::spawn(async move { app.oneshot(req).await });
tokio::time::sleep(Duration::from_millis(500)).await;
let inflight_count = tracker_clone.len();
assert!(
inflight_count > 0,
"Should have at least one in-flight request, got {}",
inflight_count
);
let buckets = tracker_clone.compute_bucket_counts();
assert!(buckets[0] > 0, "first bucket should have requests");
assert!(
*buckets.last().unwrap() > 0,
"+Inf bucket should have requests"
);
let resp = response_future.await.unwrap().unwrap();
assert_eq!(resp.status(), StatusCode::OK);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(
tracker_clone.is_empty(),
"Request should be deregistered after completion"
);
ctx.shutdown().await;
}
#[tokio::test]
async fn test_failed_request_still_deregisters() {
let ctx = AppTestContext::new(vec![MockWorkerConfig {
port: 19003,
worker_type: WorkerType::Regular,
health_status: HealthStatus::Healthy,
response_delay_ms: 0,
fail_rate: 1.0,
}])
.await;
let tracker = &ctx.app_context.inflight_tracker;
assert!(tracker.is_empty());
let app = ctx.create_app().await;
let payload = json!({
"text": "This should fail",
"stream": false
});
let req = Request::builder()
.method("POST")
.uri("/generate")
.header(CONTENT_TYPE, "application/json")
.body(Body::from(serde_json::to_string(&payload).unwrap()))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert!(tracker.is_empty());
ctx.shutdown().await;
}