diff --git a/sgl-model-gateway/src/app_context.rs b/sgl-model-gateway/src/app_context.rs index 02c56f8ab..7c1b81cfa 100644 --- a/sgl-model-gateway/src/app_context.rs +++ b/sgl-model-gateway/src/app_context.rs @@ -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>>, pub wasm_manager: Option>, pub worker_service: Arc, + pub inflight_tracker: Arc, } 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(), }) } diff --git a/sgl-model-gateway/src/middleware.rs b/sgl-model-gateway/src/middleware.rs index 5ad2c1875..84d51b321 100644 --- a/sgl-model-gateway/src/middleware.rs +++ b/sgl-model-gateway/src/middleware.rs @@ -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, +} impl HttpMetricsLayer { - pub fn new() -> Self { - Self + pub fn new(tracker: Arc) -> Self { + Self { tracker } } } @@ -620,7 +625,10 @@ impl Layer for HttpMetricsLayer { type Service = HttpMetricsMiddleware; fn layer(&self, inner: S) -> Self::Service { - HttpMetricsMiddleware { inner } + HttpMetricsMiddleware { + inner, + in_flight_request_tracker: self.tracker.clone(), + } } } @@ -628,6 +636,7 @@ impl Layer for HttpMetricsLayer { #[derive(Clone)] pub struct HttpMetricsMiddleware { inner: S, + in_flight_request_tracker: Arc, } impl Service for HttpMetricsMiddleware @@ -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); diff --git a/sgl-model-gateway/src/observability/inflight_tracker.rs b/sgl-model-gateway/src/observability/inflight_tracker.rs new file mode 100644 index 000000000..19a0639f5 --- /dev/null +++ b/sgl-model-gateway/src/observability/inflight_tracker.rs @@ -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, + next_id: AtomicU64, + sampler: OnceLock, +} + +impl InFlightRequestTracker { + pub fn new() -> Arc { + Arc::new(Self { + requests: DashMap::new(), + next_id: AtomicU64::new(0), + sampler: OnceLock::new(), + }) + } + + pub fn start_sampler(self: &Arc, 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) -> 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 = 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, + 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::>() + })); + } + + 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); + } +} diff --git a/sgl-model-gateway/src/observability/metrics.rs b/sgl-model-gateway/src/observability/metrics.rs index c4ef2ffe0..3cabf199c 100644 --- a/sgl-model-gateway/src/observability/metrics.rs +++ b/sgl-model-gateway/src/observability/metrics.rs @@ -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); } diff --git a/sgl-model-gateway/src/observability/mod.rs b/sgl-model-gateway/src/observability/mod.rs index 8a8f436d5..cd4677430 100644 --- a/sgl-model-gateway/src/observability/mod.rs +++ b/sgl-model-gateway/src/observability/mod.rs @@ -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; diff --git a/sgl-model-gateway/src/server.rs b/sgl-model-gateway/src/server.rs index f3bfe6032..1b0a7e556 100644 --- a/sgl-model-gateway/src/server.rs +++ b/sgl-model-gateway/src/server.rs @@ -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 Arc { - 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(), }) } diff --git a/sgl-model-gateway/tests/inflight_tracker_test.rs b/sgl-model-gateway/tests/inflight_tracker_test.rs new file mode 100644 index 000000000..ee39da32d --- /dev/null +++ b/sgl-model-gateway/tests/inflight_tracker_test.rs @@ -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; +}