From 46a5a37f558a5eb362fdb4bd34758331663c3d37 Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Tue, 18 Aug 2026 16:40:19 +0000 Subject: [PATCH 1/9] feat(pylon): translate Dynamo stats correlations Signed-off-by: Barry Greengus --- .../pylon-lib/src/quic_http_tunnel/backend.rs | 28 ++++++-- .../pylon-lib/src/quic_http_tunnel/core.rs | 25 +++++-- .../pylon-lib/src/quic_http_tunnel/tests.rs | 52 +++++++++++++- .../crates/pylon-lib/src/runtime_state.rs | 26 +++++++ .../src/stats/engine_stats_stream.rs | 71 ++++++++++++++++++- 5 files changed, 190 insertions(+), 12 deletions(-) diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs index c797388e4..afaf96db9 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs @@ -62,22 +62,42 @@ pub const DEFAULT_PRIORITY_CEILING: u32 = 3600; pub(crate) mod dynamo { use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; + use stargate_protocol::tunnel_contract::{HEADER_MODEL, HEADER_REQUEST_ID, HEADER_ROUTING_KEY}; + use uuid::Uuid; /// Engine priority headers pylon derives; the names stay out of the /// shared tunnel contract because only pylon speaks them. + pub(crate) const HEADER_STATS_CORRELATION_ID: &str = "x-dynamo-stats-correlation-id"; pub(crate) const HEADER_REQUEST_PRIORITY: &str = "x-dynamo-request-priority"; pub(crate) const HEADER_REQUEST_STRICT_PRIORITY: &str = "x-dynamo-request-strict-priority"; /// Denylist of engine headers pylon owns: inbound values are stripped in - /// every backend mode so pylon stays their only writer. Scoped to the - /// priority headers for now; other engine headers are tracked separately. - const STRIPPED_REQUEST_HEADERS: [&str; 2] = - [HEADER_REQUEST_PRIORITY, HEADER_REQUEST_STRICT_PRIORITY]; + /// every backend mode so pylon stays their only writer. + const STRIPPED_REQUEST_HEADERS: [&str; 5] = [ + "request-id", + "x-dynamo-request-id", + HEADER_STATS_CORRELATION_ID, + HEADER_REQUEST_PRIORITY, + HEADER_REQUEST_STRICT_PRIORITY, + ]; pub(crate) fn is_stripped_engine_header(name: &HeaderName) -> bool { STRIPPED_REQUEST_HEADERS.contains(&name.as_str()) } + /// Replace platform identity headers with an engine-local stats correlation ID. + pub(crate) fn translate_stats_correlation(upstream_headers: &mut HeaderMap) -> String { + for name in [HEADER_REQUEST_ID, HEADER_MODEL, HEADER_ROUTING_KEY] { + upstream_headers.remove(name); + } + let correlation_id = Uuid::new_v4().to_string(); + upstream_headers.insert( + HeaderName::from_static(HEADER_STATS_CORRELATION_ID), + HeaderValue::from_str(&correlation_id).expect("UUID should be a valid header value"), + ); + correlation_id + } + /// Map the platform rank (lower wins, absent = unconfigured) to the /// engine value (higher wins, read as seconds of queue head start): /// `max(0, ceiling - rank)`, with absent as the lowest value. The head diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs index b991de383..847fc39a9 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs @@ -677,9 +677,6 @@ pub(super) async fn forward_tunnel_request( } } - let priority = lifecycle - .as_ref() - .and_then(|lifecycle| lifecycle.required.priority); let response = match send_upstream_request( app, method, @@ -687,7 +684,7 @@ pub(super) async fn forward_tunnel_request( &request_headers, body_bytes, health_request, - priority, + lifecycle.as_ref(), ) .await { @@ -727,8 +724,10 @@ async fn send_upstream_request( request_headers: &HeaderMap, body_bytes: Vec, health_request: bool, - priority: Option, + lifecycle: Option<&TunnelRequestLifecycle>, ) -> Result { + let priority = lifecycle.and_then(|lifecycle| lifecycle.required.priority); + let generation = lifecycle.and_then(|lifecycle| lifecycle.generation.as_ref()); let span = if !health_request { let span = tracing::info_span!( "pylon_upstream_http_request", @@ -755,11 +754,19 @@ async fn send_upstream_request( upstream_headers.append(name, value.clone()); } } + let mut registered_stats_correlation_id = None; if !health_request { if let Some(priority) = priority { span.record("priority", priority); } if app.upstream_backend == UpstreamBackend::Dynamo { + let correlation_id = + backend::dynamo::translate_stats_correlation(&mut upstream_headers); + if let Some(generation) = generation { + app.runtime_state + .register_engine_stats_correlation(correlation_id.clone(), generation.clone()); + registered_stats_correlation_id = Some(correlation_id); + } let dynamo_priority = backend::dynamo::apply_priority_headers( priority, app.priority_ceiling, @@ -781,6 +788,14 @@ async fn send_upstream_request( .map_err(UpstreamRequestError::Send) }; let result = send.instrument(span.clone()).await; + let request_failed = match &result { + Ok(response) => !response.status().is_success(), + Err(_) => true, + }; + if request_failed && let Some(correlation_id) = registered_stats_correlation_id { + app.runtime_state + .finish_engine_stats_correlation(&correlation_id); + } match &result { Ok(response) => span.record("upstream.status", response.status().as_u16()), Err(error) => span.record("upstream.error", error.to_string()), diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs index 387704802..3d1977457 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs @@ -480,6 +480,9 @@ fn pylon_request_header_filter_strips_tunnel_headers_case_insensitively() "X-Method", "X-Path", "X-Stargate-Expected-Queue-Ms", + "Request-Id", + "X-Dynamo-Request-Id", + "X-Dynamo-Stats-Correlation-Id", "X-Dynamo-Request-Priority", "X-Dynamo-Request-Strict-Priority", ] @@ -538,6 +541,22 @@ fn pylon_dynamo_priority_headers_are_always_emitted() { assert_eq!(headers["x-dynamo-request-strict-priority"], "0"); } +#[test] +fn pylon_replaces_platform_identity_with_a_dynamo_stats_correlation() { + let mut headers = HeaderMap::new(); + headers.insert("x-request-id", "gateway-request".parse().unwrap()); + headers.insert("x-model", "gateway-model".parse().unwrap()); + headers.insert("x-routing-key", "gateway-route".parse().unwrap()); + + let correlation_id = dynamo::translate_stats_correlation(&mut headers); + + assert!(uuid::Uuid::parse_str(&correlation_id).is_ok()); + assert_eq!(headers[dynamo::HEADER_STATS_CORRELATION_ID], correlation_id); + assert!(!headers.contains_key("x-request-id")); + assert!(!headers.contains_key("x-model")); + assert!(!headers.contains_key("x-routing-key")); +} + #[test] fn pylon_trace_context_extracts_remote_parent() -> Result<()> { opentelemetry::global::set_text_map_propagator( @@ -1721,6 +1740,10 @@ fn dynamo_priority_echo_router() -> Router { }; let dynamo_priority = echo_header("x-dynamo-request-priority"); let dynamo_strict_priority = echo_header("x-dynamo-request-strict-priority"); + let stats_correlation_id = echo_header(dynamo::HEADER_STATS_CORRELATION_ID); + let platform_identity_present = ["x-request-id", "x-model", "x-routing-key"] + .into_iter() + .any(|name| req.headers().contains_key(name)); let mut sse = axum::response::Sse::new(async_stream::stream! { yield Ok::<_, std::convert::Infallible>( Event::default().data(r#"{"object":"chat.completion.chunk","choices":[{"delta":{"content":"ok"}}]}"#) @@ -1736,6 +1759,14 @@ fn dynamo_priority_echo_router() -> Router { HeaderName::from_static("x-echo-dynamo-strict-priority"), HeaderValue::from_str(&dynamo_strict_priority).unwrap(), ); + sse.headers_mut().insert( + HeaderName::from_static("x-echo-stats-correlation-id"), + HeaderValue::from_str(&stats_correlation_id).unwrap(), + ); + sse.headers_mut().insert( + HeaderName::from_static("x-saw-platform-identity"), + HeaderValue::from_static(if platform_identity_present { "true" } else { "false" }), + ); *sse.status_mut() = StatusCode::OK; sse }), @@ -1743,9 +1774,10 @@ fn dynamo_priority_echo_router() -> Router { } #[tokio::test] -async fn quic_tunnel_derives_dynamo_priority_from_x_priority() { +async fn quic_tunnel_translates_dynamo_request_headers() { let (config, _metrics) = metered_test_tunnel_config_for(dynamo_priority_echo_router()).await; let ceiling = config.forwarding.priority_ceiling; + let runtime_state = config.forwarding.runtime_state.clone(); let mut tunnel = RawTunnelTest::start(config).await; let mut headers = @@ -1754,6 +1786,11 @@ async fn quic_tunnel_derives_dynamo_priority_from_x_priority() { // Spoofed engine headers must be replaced by pylon-derived values. headers.insert("x-dynamo-request-priority", "42".parse().unwrap()); headers.insert("x-dynamo-request-strict-priority", "1".parse().unwrap()); + headers.insert( + "request-id", + uuid::Uuid::new_v4().to_string().parse().unwrap(), + ); + headers.insert("x-routing-key", "gateway-route".parse().unwrap()); tunnel .send(headers, br#"{"messages":[],"stream":true}"#) .await; @@ -1768,6 +1805,19 @@ async fn quic_tunnel_derives_dynamo_priority_from_x_priority() { (ceiling - 7).to_string() ); assert_eq!(response_headers["x-echo-dynamo-strict-priority"], "0"); + assert_eq!(response_headers["x-saw-platform-identity"], "false"); + let stats_correlation_id = response_headers["x-echo-stats-correlation-id"] + .to_str() + .unwrap(); + assert!(uuid::Uuid::parse_str(stats_correlation_id).is_ok()); + assert_eq!( + runtime_state + .engine_stats_generation(stats_correlation_id) + .as_ref() + .map(ModelGeneration::model_id), + Some("model-a") + ); + runtime_state.finish_engine_stats_correlation(stats_correlation_id); tunnel.shutdown().await; } diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs index 0cc849340..0b1e62907 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs @@ -85,6 +85,7 @@ impl ModelGeneration { pub struct PylonRuntimeState { advertised: Arc>, live_requests: LiveRequestState, + engine_stats_correlations: Arc>>, metrics: Option>, observation_tx: Option>, } @@ -168,6 +169,7 @@ impl PylonRuntimeState { models, })), live_requests: LiveRequestState::default(), + engine_stats_correlations: Arc::default(), metrics: None, observation_tx: None, } @@ -275,6 +277,9 @@ impl PylonRuntimeState { .remove(generation.model_id()) .expect("validated generation should still exist"); self.live_requests.retire_generation(generation); + self.engine_stats_correlations + .lock() + .retain(|_, owner| owner != generation); Some(retired.stats) } @@ -506,6 +511,27 @@ impl PylonRuntimeState { self.live_requests.request_generation(request_id) } + pub(crate) fn register_engine_stats_correlation( + &self, + correlation_id: String, + generation: ModelGeneration, + ) { + self.engine_stats_correlations + .lock() + .insert(correlation_id, generation); + } + + pub(crate) fn engine_stats_generation(&self, correlation_id: &str) -> Option { + self.engine_stats_correlations + .lock() + .get(correlation_id) + .cloned() + } + + pub(crate) fn finish_engine_stats_correlation(&self, correlation_id: &str) { + self.engine_stats_correlations.lock().remove(correlation_id); + } + pub(crate) fn snapshot_live_model(&self, model_id: &str) -> QueueModelSnapshot { self.current_generation(model_id) .map_or_else(QueueModelSnapshot::default, |generation| { diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs index a34319160..c903ec126 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs @@ -176,7 +176,11 @@ fn parse_stats_event( raw: RawEngineStatsEvent<'_>, observed_at: TokioInstant, ) -> Result { - let request_id = required_nonempty_string(raw.request_id, "request_id")?; + let dynamo_request_id = required_nonempty_string(raw.request_id, "request_id")?; + let request_id = match raw.correlation_id { + Some(value) => required_nonempty_string(Some(value), "correlation_id")?, + None => dynamo_request_id, + }; let model_id = required_nonempty_string(raw.model, "model")?; let tokens_processed = optional_u64(raw.tokens_processed, "tokens_processed")?; let tokens_generated = optional_u64(raw.tokens_generated, "tokens_generated")?; @@ -209,6 +213,7 @@ struct RawEngineStatsEvent<'a> { version: Option>, event_type: Option>, request_id: Option>, + correlation_id: Option>, model: Option>, tokens_processed: Option>, tokens_generated: Option>, @@ -243,6 +248,7 @@ impl<'de> Visitor<'de> for RawEngineStatsEventVisitor { "v" => event.version = Some(map.next_value()?), "type" => event.event_type = Some(map.next_value()?), "request_id" => event.request_id = Some(map.next_value()?), + "correlation_id" => event.correlation_id = Some(map.next_value()?), "model" => event.model = Some(map.next_value()?), "tokens_processed" => event.tokens_processed = Some(map.next_value()?), "tokens_generated" => event.tokens_generated = Some(map.next_value()?), @@ -598,12 +604,24 @@ async fn emit_engine_stats_event( } match update { Some(mut update) => { - update.generation = generated_request_generation(&update.request_id, &update.model_id) + let platform_generation = config.runtime_state.as_ref().and_then(|runtime_state| { + runtime_state.engine_stats_generation(&update.request_id) + }); + if let Some(generation) = &platform_generation { + update.model_id = generation.model_id().to_string(); + } + update.generation = platform_generation + .or_else(|| generated_request_generation(&update.request_id, &update.model_id)) .or_else(|| { config.runtime_state.as_ref().and_then(|runtime_state| { runtime_state.request_generation(&update.request_id) }) }); + if update.finished + && let Some(runtime_state) = &config.runtime_state + { + runtime_state.finish_engine_stats_correlation(&update.request_id); + } send_stats_update( stats_update_tx, StatsAggregatorUpdate::RequestCounters(update), @@ -826,6 +844,47 @@ mod tests { assert_eq!(update.generation, Some(generation)); } + #[tokio::test] + async fn dynamo_stats_correlation_is_translated_to_the_platform_generation() { + let runtime_state = PylonRuntimeState::new( + stargate_proto::pb::InferenceServerStatus::Active, + &["platform-model".to_string()], + ); + let generation = runtime_state + .current_generation("platform-model") + .expect("test generation should exist"); + runtime_state.register_engine_stats_correlation( + "dynamo-correlation".to_string(), + generation.clone(), + ); + let config = EngineStatsStreamConfig { + runtime_state: Some(runtime_state.clone()), + ..EngineStatsStreamConfig::default() + }; + + let processed = process_lines( + &config, + [r#"{"v":1,"type":"stats","request_id":"dynamo-request","correlation_id":"dynamo-correlation","model":"dynamo-model","tokens_processed":64,"finished":true} +"#], + ) + .await; + let StatsAggregatorUpdate::RequestCounters(update) = processed + .updates + .try_recv() + .expect("engine event should enter the stats pipeline") + else { + panic!("expected request counters update"); + }; + + assert_eq!(update.request_id, "dynamo-correlation"); + assert_eq!(update.model_id, "platform-model"); + assert_eq!(update.generation, Some(generation)); + assert_eq!( + runtime_state.engine_stats_generation("dynamo-correlation"), + None + ); + } + #[test] fn rejects_invalid_engine_stats_events() { assert!(matches!( @@ -870,6 +929,14 @@ mod tests { br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","finished":"true"}"#.as_slice(), "finished", ), + ( + br#"{"v":1,"type":"stats","request_id":"req-1","correlation_id":1,"model":"llama","finished":true}"#.as_slice(), + "correlation_id", + ), + ( + br#"{"v":1,"type":"stats","request_id":"req-1","correlation_id":" ","model":"llama","finished":true}"#.as_slice(), + "correlation_id", + ), ] { assert!(matches!( parse(json).unwrap_err(), From 7f99cbec8dcda8ef6f1d355e048bc2df952ac597 Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Tue, 18 Aug 2026 17:50:25 +0000 Subject: [PATCH 2/9] refactor(pylon): use canonical Dynamo request IDs Signed-off-by: Barry Greengus --- .../stargate/crates/pylon-lib/src/bringup.rs | 7 +++ .../crates/pylon-lib/src/bringup/upstream.rs | 1 + .../pylon-lib/src/quic_http_tunnel/backend.rs | 30 +++++---- .../pylon-lib/src/quic_http_tunnel/core.rs | 21 ++----- .../pylon-lib/src/quic_http_tunnel/tests.rs | 56 +++++++++-------- .../crates/pylon-lib/src/runtime_state.rs | 26 -------- .../src/stats/engine_stats_stream.rs | 61 +++++++------------ 7 files changed, 83 insertions(+), 119 deletions(-) diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs index 0e5492484..a5e3d4a68 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs @@ -928,6 +928,13 @@ mod tests { .get(HEADER_REQUEST_ID) .and_then(|value| value.to_str().ok()) { + assert_eq!( + headers + .get("request-id") + .and_then(|value| value.to_str().ok()), + Some(request_id), + "Pylon-generated requests must use one canonical upstream ID" + ); request_ids.lock().await.push(request_id.to_string()); } let prompt = request diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs index 0671039fc..93007d2e6 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs @@ -116,6 +116,7 @@ pub(super) async fn send_completion_request( "/v1/chat/completions", )) .header(HEADER_REQUEST_ID, &request_id) + .header("request-id", &request_id) .header(HEADER_MODEL, model_id) .header(HEADER_INPUT_TOKENS, input_tokens.to_string()) .json(request); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs index afaf96db9..f3110aad1 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs @@ -63,20 +63,18 @@ pub const DEFAULT_PRIORITY_CEILING: u32 = 3600; pub(crate) mod dynamo { use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; use stargate_protocol::tunnel_contract::{HEADER_MODEL, HEADER_REQUEST_ID, HEADER_ROUTING_KEY}; - use uuid::Uuid; /// Engine priority headers pylon derives; the names stay out of the /// shared tunnel contract because only pylon speaks them. - pub(crate) const HEADER_STATS_CORRELATION_ID: &str = "x-dynamo-stats-correlation-id"; + pub(crate) const HEADER_DYNAMO_REQUEST_ID: &str = "request-id"; pub(crate) const HEADER_REQUEST_PRIORITY: &str = "x-dynamo-request-priority"; pub(crate) const HEADER_REQUEST_STRICT_PRIORITY: &str = "x-dynamo-request-strict-priority"; /// Denylist of engine headers pylon owns: inbound values are stripped in /// every backend mode so pylon stays their only writer. - const STRIPPED_REQUEST_HEADERS: [&str; 5] = [ - "request-id", + const STRIPPED_REQUEST_HEADERS: [&str; 4] = [ + HEADER_DYNAMO_REQUEST_ID, "x-dynamo-request-id", - HEADER_STATS_CORRELATION_ID, HEADER_REQUEST_PRIORITY, HEADER_REQUEST_STRICT_PRIORITY, ]; @@ -85,17 +83,25 @@ pub(crate) mod dynamo { STRIPPED_REQUEST_HEADERS.contains(&name.as_str()) } - /// Replace platform identity headers with an engine-local stats correlation ID. - pub(crate) fn translate_stats_correlation(upstream_headers: &mut HeaderMap) -> String { - for name in [HEADER_REQUEST_ID, HEADER_MODEL, HEADER_ROUTING_KEY] { + /// Translate the validated platform request ID into Dynamo's canonical ID. + pub(crate) fn apply_request_id(request_id: &str, upstream_headers: &mut HeaderMap) { + for name in [HEADER_DYNAMO_REQUEST_ID, "x-dynamo-request-id"] { + upstream_headers.remove(name); + } + for name in [HEADER_MODEL, HEADER_ROUTING_KEY] { upstream_headers.remove(name); } - let correlation_id = Uuid::new_v4().to_string(); upstream_headers.insert( - HeaderName::from_static(HEADER_STATS_CORRELATION_ID), - HeaderValue::from_str(&correlation_id).expect("UUID should be a valid header value"), + HeaderName::from_static(HEADER_DYNAMO_REQUEST_ID), + HeaderValue::from_str(request_id) + .expect("validated x-request-id should be a valid header value"), + ); + debug_assert_eq!( + upstream_headers + .get(HEADER_REQUEST_ID) + .and_then(|value| value.to_str().ok()), + Some(request_id) ); - correlation_id } /// Map the platform rank (lower wins, absent = unconfigured) to the diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs index 847fc39a9..12cd1b277 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs @@ -727,7 +727,6 @@ async fn send_upstream_request( lifecycle: Option<&TunnelRequestLifecycle>, ) -> Result { let priority = lifecycle.and_then(|lifecycle| lifecycle.required.priority); - let generation = lifecycle.and_then(|lifecycle| lifecycle.generation.as_ref()); let span = if !health_request { let span = tracing::info_span!( "pylon_upstream_http_request", @@ -754,18 +753,16 @@ async fn send_upstream_request( upstream_headers.append(name, value.clone()); } } - let mut registered_stats_correlation_id = None; if !health_request { if let Some(priority) = priority { span.record("priority", priority); } if app.upstream_backend == UpstreamBackend::Dynamo { - let correlation_id = - backend::dynamo::translate_stats_correlation(&mut upstream_headers); - if let Some(generation) = generation { - app.runtime_state - .register_engine_stats_correlation(correlation_id.clone(), generation.clone()); - registered_stats_correlation_id = Some(correlation_id); + if let Some(lifecycle) = lifecycle { + backend::dynamo::apply_request_id( + &lifecycle.required.request_id, + &mut upstream_headers, + ); } let dynamo_priority = backend::dynamo::apply_priority_headers( priority, @@ -788,14 +785,6 @@ async fn send_upstream_request( .map_err(UpstreamRequestError::Send) }; let result = send.instrument(span.clone()).await; - let request_failed = match &result { - Ok(response) => !response.status().is_success(), - Err(_) => true, - }; - if request_failed && let Some(correlation_id) = registered_stats_correlation_id { - app.runtime_state - .finish_engine_stats_correlation(&correlation_id); - } match &result { Ok(response) => span.record("upstream.status", response.status().as_u16()), Err(error) => span.record("upstream.error", error.to_string()), diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs index 3d1977457..ed7b36206 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs @@ -482,7 +482,6 @@ fn pylon_request_header_filter_strips_tunnel_headers_case_insensitively() "X-Stargate-Expected-Queue-Ms", "Request-Id", "X-Dynamo-Request-Id", - "X-Dynamo-Stats-Correlation-Id", "X-Dynamo-Request-Priority", "X-Dynamo-Request-Strict-Priority", ] @@ -542,17 +541,19 @@ fn pylon_dynamo_priority_headers_are_always_emitted() { } #[test] -fn pylon_replaces_platform_identity_with_a_dynamo_stats_correlation() { +fn pylon_translates_platform_request_id_to_dynamo_request_id() { let mut headers = HeaderMap::new(); headers.insert("x-request-id", "gateway-request".parse().unwrap()); headers.insert("x-model", "gateway-model".parse().unwrap()); headers.insert("x-routing-key", "gateway-route".parse().unwrap()); + headers.insert("request-id", "spoofed-request".parse().unwrap()); + headers.insert("x-dynamo-request-id", "spoofed-legacy".parse().unwrap()); - let correlation_id = dynamo::translate_stats_correlation(&mut headers); + dynamo::apply_request_id("gateway-request", &mut headers); - assert!(uuid::Uuid::parse_str(&correlation_id).is_ok()); - assert_eq!(headers[dynamo::HEADER_STATS_CORRELATION_ID], correlation_id); - assert!(!headers.contains_key("x-request-id")); + assert_eq!(headers["x-request-id"], "gateway-request"); + assert_eq!(headers["request-id"], "gateway-request"); + assert!(!headers.contains_key("x-dynamo-request-id")); assert!(!headers.contains_key("x-model")); assert!(!headers.contains_key("x-routing-key")); } @@ -1071,6 +1072,7 @@ async fn http3_direct_tunnel_accepts_responses_request_to_upstream() { ); let mut config = test_tunnel_config_for(app).await; config.tunnel_protocol = TunnelTransportProtocol::Http3; + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; let tunnel = start_quic_http_tunnel(config).await.unwrap(); let mut headers = HeaderMap::new(); headers.insert("x-request-id", "req-h3-direct".parse().unwrap()); @@ -1669,6 +1671,7 @@ async fn quic_tunnel_forwards_to_http_backend() { }), ); let (mut config, _metrics) = metered_test_tunnel_config_for(app).await; + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; config.forwarding.retry.upstream_retry_header = HeaderName::from_static("x-vendor-retryable"); let mut tunnel = RawTunnelTest::start(config).await; @@ -1740,8 +1743,9 @@ fn dynamo_priority_echo_router() -> Router { }; let dynamo_priority = echo_header("x-dynamo-request-priority"); let dynamo_strict_priority = echo_header("x-dynamo-request-strict-priority"); - let stats_correlation_id = echo_header(dynamo::HEADER_STATS_CORRELATION_ID); - let platform_identity_present = ["x-request-id", "x-model", "x-routing-key"] + let x_request_id = echo_header("x-request-id"); + let dynamo_request_id = echo_header(dynamo::HEADER_DYNAMO_REQUEST_ID); + let platform_routing_identity_present = ["x-model", "x-routing-key"] .into_iter() .any(|name| req.headers().contains_key(name)); let mut sse = axum::response::Sse::new(async_stream::stream! { @@ -1760,12 +1764,16 @@ fn dynamo_priority_echo_router() -> Router { HeaderValue::from_str(&dynamo_strict_priority).unwrap(), ); sse.headers_mut().insert( - HeaderName::from_static("x-echo-stats-correlation-id"), - HeaderValue::from_str(&stats_correlation_id).unwrap(), + HeaderName::from_static("x-echo-request-id"), + HeaderValue::from_str(&x_request_id).unwrap(), + ); + sse.headers_mut().insert( + HeaderName::from_static("x-echo-dynamo-request-id"), + HeaderValue::from_str(&dynamo_request_id).unwrap(), ); sse.headers_mut().insert( - HeaderName::from_static("x-saw-platform-identity"), - HeaderValue::from_static(if platform_identity_present { "true" } else { "false" }), + HeaderName::from_static("x-saw-platform-routing-identity"), + HeaderValue::from_static(if platform_routing_identity_present { "true" } else { "false" }), ); *sse.status_mut() = StatusCode::OK; sse @@ -1777,7 +1785,6 @@ fn dynamo_priority_echo_router() -> Router { async fn quic_tunnel_translates_dynamo_request_headers() { let (config, _metrics) = metered_test_tunnel_config_for(dynamo_priority_echo_router()).await; let ceiling = config.forwarding.priority_ceiling; - let runtime_state = config.forwarding.runtime_state.clone(); let mut tunnel = RawTunnelTest::start(config).await; let mut headers = @@ -1790,6 +1797,10 @@ async fn quic_tunnel_translates_dynamo_request_headers() { "request-id", uuid::Uuid::new_v4().to_string().parse().unwrap(), ); + headers.insert( + "x-dynamo-request-id", + uuid::Uuid::new_v4().to_string().parse().unwrap(), + ); headers.insert("x-routing-key", "gateway-route".parse().unwrap()); tunnel .send(headers, br#"{"messages":[],"stream":true}"#) @@ -1805,19 +1816,9 @@ async fn quic_tunnel_translates_dynamo_request_headers() { (ceiling - 7).to_string() ); assert_eq!(response_headers["x-echo-dynamo-strict-priority"], "0"); - assert_eq!(response_headers["x-saw-platform-identity"], "false"); - let stats_correlation_id = response_headers["x-echo-stats-correlation-id"] - .to_str() - .unwrap(); - assert!(uuid::Uuid::parse_str(stats_correlation_id).is_ok()); - assert_eq!( - runtime_state - .engine_stats_generation(stats_correlation_id) - .as_ref() - .map(ModelGeneration::model_id), - Some("model-a") - ); - runtime_state.finish_engine_stats_correlation(stats_correlation_id); + assert_eq!(response_headers["x-echo-request-id"], "req-dynamo-1"); + assert_eq!(response_headers["x-echo-dynamo-request-id"], "req-dynamo-1"); + assert_eq!(response_headers["x-saw-platform-routing-identity"], "false"); tunnel.shutdown().await; } @@ -1901,6 +1902,8 @@ async fn quic_tunnel_passthrough_backend_strips_but_derives_nothing() { assert_eq!(response_headers["x-echo-dynamo-priority"], "absent"); // Stripping inbound engine priority headers is not gated by the backend. assert_eq!(response_headers["x-echo-dynamo-strict-priority"], "absent"); + assert_eq!(response_headers["x-echo-request-id"], "req-dynamo-3"); + assert_eq!(response_headers["x-echo-dynamo-request-id"], "absent"); tunnel.shutdown().await; } @@ -2754,6 +2757,7 @@ async fn assert_direct_embeddings_case(case: DirectEmbeddingsCase) { let (runtime_state, rx) = observed_runtime(16); let mut config = test_tunnel_config_for(app).await; config.tunnel_protocol = case.protocol; + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; config.forwarding.runtime_state = runtime_state; let tunnel = start_quic_http_tunnel(config).await.unwrap(); let client = DirectTunnelClient::connect(case.protocol, tunnel.listen_addr()).await; diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs index 0b1e62907..0cc849340 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs @@ -85,7 +85,6 @@ impl ModelGeneration { pub struct PylonRuntimeState { advertised: Arc>, live_requests: LiveRequestState, - engine_stats_correlations: Arc>>, metrics: Option>, observation_tx: Option>, } @@ -169,7 +168,6 @@ impl PylonRuntimeState { models, })), live_requests: LiveRequestState::default(), - engine_stats_correlations: Arc::default(), metrics: None, observation_tx: None, } @@ -277,9 +275,6 @@ impl PylonRuntimeState { .remove(generation.model_id()) .expect("validated generation should still exist"); self.live_requests.retire_generation(generation); - self.engine_stats_correlations - .lock() - .retain(|_, owner| owner != generation); Some(retired.stats) } @@ -511,27 +506,6 @@ impl PylonRuntimeState { self.live_requests.request_generation(request_id) } - pub(crate) fn register_engine_stats_correlation( - &self, - correlation_id: String, - generation: ModelGeneration, - ) { - self.engine_stats_correlations - .lock() - .insert(correlation_id, generation); - } - - pub(crate) fn engine_stats_generation(&self, correlation_id: &str) -> Option { - self.engine_stats_correlations - .lock() - .get(correlation_id) - .cloned() - } - - pub(crate) fn finish_engine_stats_correlation(&self, correlation_id: &str) { - self.engine_stats_correlations.lock().remove(correlation_id); - } - pub(crate) fn snapshot_live_model(&self, model_id: &str) -> QueueModelSnapshot { self.current_generation(model_id) .map_or_else(QueueModelSnapshot::default, |generation| { diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs index c903ec126..607f68640 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs @@ -176,11 +176,7 @@ fn parse_stats_event( raw: RawEngineStatsEvent<'_>, observed_at: TokioInstant, ) -> Result { - let dynamo_request_id = required_nonempty_string(raw.request_id, "request_id")?; - let request_id = match raw.correlation_id { - Some(value) => required_nonempty_string(Some(value), "correlation_id")?, - None => dynamo_request_id, - }; + let request_id = required_nonempty_string(raw.request_id, "request_id")?; let model_id = required_nonempty_string(raw.model, "model")?; let tokens_processed = optional_u64(raw.tokens_processed, "tokens_processed")?; let tokens_generated = optional_u64(raw.tokens_generated, "tokens_generated")?; @@ -213,7 +209,6 @@ struct RawEngineStatsEvent<'a> { version: Option>, event_type: Option>, request_id: Option>, - correlation_id: Option>, model: Option>, tokens_processed: Option>, tokens_generated: Option>, @@ -248,7 +243,6 @@ impl<'de> Visitor<'de> for RawEngineStatsEventVisitor { "v" => event.version = Some(map.next_value()?), "type" => event.event_type = Some(map.next_value()?), "request_id" => event.request_id = Some(map.next_value()?), - "correlation_id" => event.correlation_id = Some(map.next_value()?), "model" => event.model = Some(map.next_value()?), "tokens_processed" => event.tokens_processed = Some(map.next_value()?), "tokens_generated" => event.tokens_generated = Some(map.next_value()?), @@ -604,24 +598,12 @@ async fn emit_engine_stats_event( } match update { Some(mut update) => { - let platform_generation = config.runtime_state.as_ref().and_then(|runtime_state| { - runtime_state.engine_stats_generation(&update.request_id) - }); - if let Some(generation) = &platform_generation { - update.model_id = generation.model_id().to_string(); - } - update.generation = platform_generation - .or_else(|| generated_request_generation(&update.request_id, &update.model_id)) + update.generation = generated_request_generation(&update.request_id, &update.model_id) .or_else(|| { config.runtime_state.as_ref().and_then(|runtime_state| { runtime_state.request_generation(&update.request_id) }) }); - if update.finished - && let Some(runtime_state) = &config.runtime_state - { - runtime_state.finish_engine_stats_correlation(&update.request_id); - } send_stats_update( stats_update_tx, StatsAggregatorUpdate::RequestCounters(update), @@ -700,6 +682,9 @@ mod tests { use tokio::net::TcpListener; use crate::generated_request_id::{GeneratedRequestKind, next_generated_request_id}; + use crate::request_observer::{ + RequestObservationEndpoint, RequiredTunnelHeaders, TunnelRequestObserver, + }; fn parse(line: &[u8]) -> Result { parse_engine_stats_line(line, TokioInstant::now()) @@ -845,7 +830,7 @@ mod tests { } #[tokio::test] - async fn dynamo_stats_correlation_is_translated_to_the_platform_generation() { + async fn platform_request_id_routes_engine_stats_to_the_live_generation() { let runtime_state = PylonRuntimeState::new( stargate_proto::pb::InferenceServerStatus::Active, &["platform-model".to_string()], @@ -853,18 +838,27 @@ mod tests { let generation = runtime_state .current_generation("platform-model") .expect("test generation should exist"); - runtime_state.register_engine_stats_correlation( - "dynamo-correlation".to_string(), - generation.clone(), + let observer = TunnelRequestObserver::accepted( + RequestObservationEndpoint::ChatCompletions, + RequiredTunnelHeaders { + request_id: "gateway-request".to_string(), + routing_key: None, + model_id: "platform-model".to_string(), + priority: None, + input_tokens: 64, + accepted_at: std::time::Instant::now(), + }, + Some(generation.clone()), + runtime_state.clone(), ); let config = EngineStatsStreamConfig { - runtime_state: Some(runtime_state.clone()), + runtime_state: Some(runtime_state), ..EngineStatsStreamConfig::default() }; let processed = process_lines( &config, - [r#"{"v":1,"type":"stats","request_id":"dynamo-request","correlation_id":"dynamo-correlation","model":"dynamo-model","tokens_processed":64,"finished":true} + [r#"{"v":1,"type":"stats","request_id":"gateway-request","model":"platform-model","tokens_processed":64,"finished":true} "#], ) .await; @@ -876,13 +870,10 @@ mod tests { panic!("expected request counters update"); }; - assert_eq!(update.request_id, "dynamo-correlation"); + assert_eq!(update.request_id, "gateway-request"); assert_eq!(update.model_id, "platform-model"); assert_eq!(update.generation, Some(generation)); - assert_eq!( - runtime_state.engine_stats_generation("dynamo-correlation"), - None - ); + drop(observer); } #[test] @@ -929,14 +920,6 @@ mod tests { br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","finished":"true"}"#.as_slice(), "finished", ), - ( - br#"{"v":1,"type":"stats","request_id":"req-1","correlation_id":1,"model":"llama","finished":true}"#.as_slice(), - "correlation_id", - ), - ( - br#"{"v":1,"type":"stats","request_id":"req-1","correlation_id":" ","model":"llama","finished":true}"#.as_slice(), - "correlation_id", - ), ] { assert!(matches!( parse(json).unwrap_err(), From a73f099b1680f8c3dc6f34d40987129408edfb56 Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Tue, 18 Aug 2026 23:18:50 +0000 Subject: [PATCH 3/9] feat(pylon): consume canonical Dynamo KV stats Signed-off-by: Barry Greengus --- .../stargate/crates/mock-dynamo/src/main.rs | 11 + .../stargate/crates/mock-dynamo/src/openai.rs | 55 ++- .../crates/mock-dynamo/src/test_control.rs | 16 + .../stargate/crates/mock-dynamo/src/tests.rs | 18 + .../crates/proto/proto/stargate.proto | 12 + .../rust/stargate/crates/pylon-lib/src/lib.rs | 4 +- .../crates/pylon-lib/src/model_lifecycle.rs | 6 +- .../pylon-lib/src/registration/tests.rs | 17 +- .../crates/pylon-lib/src/runtime_state.rs | 18 + .../crates/pylon-lib/src/stats/aggregator.rs | 56 ++- .../crates/pylon-lib/src/stats/collector.rs | 217 ++++++---- .../pylon-lib/src/stats/kv_stats_stream.rs | 379 ++++++++++++++++++ .../crates/pylon-lib/src/stats/mod.rs | 1 + .../crates/pylon-lib/src/stats/projection.rs | 49 ++- .../rust/stargate/crates/pylon/src/main.rs | 16 +- .../rust/stargate/crates/pylon/src/startup.rs | 3 +- .../stargate/crates/stargate/src/metrics.rs | 20 + .../src/routing_state/cluster_snapshots.rs | 81 +++- .../stargate/src/routing_state/clusters.rs | 33 +- .../stargate/src/routing_state/tests.rs | 76 ++++ 20 files changed, 973 insertions(+), 115 deletions(-) create mode 100644 src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs index d0dcf49fd..654e4a99a 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs @@ -22,6 +22,7 @@ mod test_control; mod timing; use std::sync::Arc; +use std::sync::atomic::AtomicBool; use std::time::Duration; use anyhow::Result; @@ -84,6 +85,7 @@ struct AppState { health_delay: Duration, kv_cache: Arc>, stats_events: broadcast::Sender, + kv_stats_enabled: Arc, test_control: test_control::TestControlState, } @@ -116,6 +118,7 @@ async fn main() -> Result<()> { args.kv_cache_capacity_tokens, ))), stats_events, + kv_stats_enabled: Arc::new(AtomicBool::new(true)), test_control: test_control::TestControlState::with_discovered_models([args.model_name]), }; @@ -125,12 +128,20 @@ async fn main() -> Result<()> { .route("/v1/responses", post(openai::responses)) .route("/v1/embeddings", post(openai::embeddings)) .route("/pylon/v1/stats/stream", get(stats_stream::stats_stream)) + .route( + "/v1/kv-cache/stats/stream", + get(openai::kv_cache_stats_stream), + ) .route("/kv-cache/stats", get(openai::kv_cache_stats)) .route( "/test-control/models/{model}", put(test_control::update_model_test_control), ) .route("/test-control", get(test_control::test_control_snapshot)) + .route( + "/test-control/kv-stats", + put(test_control::update_kv_stats_test_control), + ) .route( "/test-control/discovery-models", put(test_control::replace_discovery_models), diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs index 7604e2307..5aab886de 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs @@ -14,11 +14,14 @@ // limitations under the License. use axum::Json; +use axum::body::{Body, Bytes}; use axum::extract::State; -use axum::http::{HeaderMap, StatusCode}; +use axum::http::{HeaderMap, HeaderValue, StatusCode, header}; use axum::response::sse::{Event, KeepAlive, Sse}; use axum::response::{IntoResponse, Response}; use serde::{Deserialize, Serialize}; +use std::convert::Infallible; +use std::sync::atomic::Ordering; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use tokio::sync::OwnedSemaphorePermit; use tracing::info; @@ -389,6 +392,52 @@ pub(crate) async fn kv_cache_stats(State(state): State) -> Json) -> Response { + if !state.kv_stats_enabled.load(Ordering::Relaxed) { + return StatusCode::SERVICE_UNAVAILABLE.into_response(); + } + let stream = async_stream::stream! { + let mut interval = tokio::time::interval(Duration::from_secs(1)); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut snapshot_id = 1_u64; + loop { + interval.tick().await; + if !state.kv_stats_enabled.load(Ordering::Relaxed) { + break; + } + let stats = state.kv_cache.lock().await.stats(&state.model_name); + let snapshot = serde_json::json!({ + "v": 1, + "type": "kv_stats_snapshot", + "snapshot_id": snapshot_id, + "observed_at_unix_ms": unix_millis(), + "models": [{ + "model": stats.model, + "aliases": [], + "routing_cache": { + "role": "aggregated", + "capacity_tokens": stats.kv_cache_capacity_tokens, + "used_tokens": stats.kv_cache_used_tokens, + "free_tokens": stats.kv_cache_free_tokens, + "complete": true + }, + "pools": [] + }] + }); + snapshot_id = snapshot_id.saturating_add(1); + let mut line = serde_json::to_vec(&snapshot).expect("mock KV stats serialize"); + line.push(b'\n'); + yield Ok::(Bytes::from(line)); + } + }; + let mut response = Response::new(Body::from_stream(stream)); + response.headers_mut().insert( + header::CONTENT_TYPE, + HeaderValue::from_static("application/x-ndjson"), + ); + response +} + impl AppState { async fn process_input_with_cache( &self, @@ -618,3 +667,7 @@ fn rand_id() -> String { fn current_unix_timestamp() -> u64 { time_since_epoch().as_secs() } + +fn unix_millis() -> u64 { + u64::try_from(time_since_epoch().as_millis()).unwrap_or(u64::MAX) +} diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs index f77a23bfb..31f6912ea 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs @@ -15,6 +15,7 @@ use std::collections::{BTreeMap, BTreeSet}; use std::sync::Arc; +use std::sync::atomic::Ordering; use axum::Json; use axum::extract::{Path, State}; @@ -252,6 +253,21 @@ pub(crate) async fn test_control_snapshot( Json(state.test_control.snapshot().await) } +#[derive(Debug, Clone, Deserialize)] +pub(crate) struct KvStatsTestControlUpdate { + pub(crate) enabled: bool, +} + +pub(crate) async fn update_kv_stats_test_control( + State(state): State, + Json(update): Json, +) -> axum::http::StatusCode { + state + .kv_stats_enabled + .store(update.enabled, Ordering::Relaxed); + axum::http::StatusCode::NO_CONTENT +} + pub(crate) fn request_class(headers: &HeaderMap) -> TestRequestClass { match headers .get("x-request-id") diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs index d4ba751a0..0126d54ce 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs @@ -53,10 +53,28 @@ fn test_state() -> AppState { health_delay: Duration::ZERO, kv_cache: Arc::new(Mutex::new(KvCacheState::new(0))), stats_events: test_stats_events(), + kv_stats_enabled: Arc::new(AtomicBool::new(true)), test_control: TestControlState::with_discovered_models(["dummy-model".to_string()]), } } +#[tokio::test] +async fn kv_stats_test_control_does_not_disable_health() { + let state = test_state(); + let status = update_kv_stats_test_control( + State(state.clone()), + Json(KvStatsTestControlUpdate { enabled: false }), + ) + .await; + + assert_eq!(status, axum::http::StatusCode::NO_CONTENT); + assert_eq!( + kv_cache_stats_stream(State(state.clone())).await.status(), + axum::http::StatusCode::SERVICE_UNAVAILABLE + ); + assert_eq!(health(State(state)).await, "ok"); +} + async fn spawn_test_app(app: Router) -> (std::net::SocketAddr, tokio::task::JoinHandle<()>) { let listener = TcpListener::bind("127.0.0.1:0") .await diff --git a/src/libraries/rust/stargate/crates/proto/proto/stargate.proto b/src/libraries/rust/stargate/crates/proto/proto/stargate.proto index ef4fe5d1f..b46a41f6a 100644 --- a/src/libraries/rust/stargate/crates/proto/proto/stargate.proto +++ b/src/libraries/rust/stargate/crates/proto/proto/stargate.proto @@ -126,6 +126,18 @@ message ModelStats { // Sticky source labels for stat observations seen by this backend since the // model metrics state was initialized. repeated string stats_sources = 17; + // Absolute KV cache snapshot reported by the serving frontend. In shared + // clusters Stargate selects the freshest complete snapshot; it never sums + // replicated frontend observations. + optional KvCacheStats kv_cache = 18; +} + +message KvCacheStats { + uint64 capacity_tokens = 1; + uint64 used_tokens = 2; + uint64 free_tokens = 3; + uint64 source_observed_at_unix_ms = 4; + bool complete = 5; } enum InferenceServerStatus { diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs index 52481d3f9..2365f7153 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs @@ -57,7 +57,9 @@ pub use request_observer::{ RequestObservation, RequestObservationEndpoint, RequestObservationState, }; pub use request_quality_monitor::RequestQualityMonitorConfig; -pub use runtime_state::{CurrentModelStats, PylonRuntimeState, RequestObservationEvent}; +pub use runtime_state::{ + CurrentKvCacheStats, CurrentModelStats, PylonRuntimeState, RequestObservationEvent, +}; pub use stats::{ EngineStatsStreamConfig, EngineStatsStreamHandle, EngineStatsStreamMode, MetricsServerHandle, PylonMetrics, RequestCounterUpdate, RequestCounterUpdateInput, StatsAggregatorUpdate, diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs index 2af6f5c30..cd7ceedf9 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs @@ -1275,10 +1275,10 @@ mod tests { let kv_cache = BlockingKvCacheServer::spawn().await; let stats_config = StatsCollectorConfig { kv_cache_stats_url: Some(kv_cache.url()), - // Poll quickly and keep the request open so retirement parks on + // Reconnect quickly and keep the request open so retirement parks on // stats cleanup after the canary ordering point under test. - kv_cache_poll_interval: Duration::from_millis(1), - kv_cache_request_timeout: Duration::from_secs(60), + kv_cache_reconnect_interval: Duration::from_millis(1), + kv_cache_connect_timeout: Duration::from_secs(60), ..StatsCollectorConfig::default() }; let (runtime_state, observations) = PylonRuntimeState::observed( diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs index 5af16b3c3..9f0df65e5 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs @@ -40,7 +40,9 @@ use tower::util::MapRequestLayer; use crate::quic_http_tunnel::{TunnelError, TunnelForwardingConfig}; use crate::request_quality_monitor::RequestQualityMonitorConfig; -use crate::runtime_state::{CurrentModelStats, PylonRuntimeState, gated_model_status}; +use crate::runtime_state::{ + CurrentKvCacheStats, CurrentModelStats, PylonRuntimeState, gated_model_status, +}; use crate::stats::PylonMetrics; use super::discovery::*; @@ -921,6 +923,12 @@ fn runtime_snapshot_forwards_bootstrap_and_collected_stats_exactly() { kv_cache_capacity_tokens: 7, kv_cache_used_tokens: 8, kv_cache_free_tokens: 9, + kv_cache: Some(CurrentKvCacheStats { + capacity_tokens: 7, + used_tokens: 3, + free_tokens: 4, + source_observed_at_unix_ms: 16, + }), num_running_queries: 10, max_engine_concurrency: Some(11), total_query_input_size: 12, @@ -940,6 +948,13 @@ fn runtime_snapshot_forwards_bootstrap_and_collected_stats_exactly() { let stats = model.stats.as_ref().expect("stats should be present"); assert_eq!(stats.last_mean_input_tps, 3.5); assert_eq!(stats.output_tps, 2.5); + assert_eq!( + stats + .kv_cache + .as_ref() + .map(|stats| stats.source_observed_at_unix_ms), + Some(16) + ); assert_eq!( stats.queue_time_estimate_ms_by_priority, queue_time_estimate_ms_by_priority diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs index 0cc849340..9070b065f 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs @@ -47,6 +47,7 @@ pub struct CurrentModelStats { pub kv_cache_capacity_tokens: u64, pub kv_cache_used_tokens: u64, pub kv_cache_free_tokens: u64, + pub kv_cache: Option, pub num_running_queries: u64, pub max_engine_concurrency: Option, pub total_query_input_size: u64, @@ -58,6 +59,14 @@ pub struct CurrentModelStats { pub stats_sources: Vec, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct CurrentKvCacheStats { + pub capacity_tokens: u64, + pub used_tokens: u64, + pub free_tokens: u64, + pub source_observed_at_unix_ms: u64, +} + #[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] pub(crate) struct ModelGeneration { model_id: String, @@ -373,6 +382,15 @@ impl PylonRuntimeState { kv_cache_capacity_tokens: stats.kv_cache_capacity_tokens, kv_cache_used_tokens: stats.kv_cache_used_tokens, kv_cache_free_tokens: stats.kv_cache_free_tokens, + kv_cache: stats.kv_cache.as_ref().map(|kv_cache| { + stargate_proto::pb::KvCacheStats { + capacity_tokens: kv_cache.capacity_tokens, + used_tokens: kv_cache.used_tokens, + free_tokens: kv_cache.free_tokens, + source_observed_at_unix_ms: kv_cache.source_observed_at_unix_ms, + complete: true, + } + }), num_running_queries: stats.num_running_queries, max_engine_concurrency: stats.max_engine_concurrency.unwrap_or_default(), total_query_input_size: stats.total_query_input_size, diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs index c6c680883..e24285121 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs @@ -16,13 +16,12 @@ use std::collections::{HashMap, VecDeque}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; -use serde::Deserialize; use stargate_protocol::common::valid_last_mean_input_tps; use tokio::time::Instant as TokioInstant; use crate::generated_request_id::{GeneratedRequestKind, generated_request_kind}; use crate::runtime_state::ModelGeneration; -use crate::{CurrentModelStats, PylonRuntimeState}; +use crate::{CurrentKvCacheStats, CurrentModelStats, PylonRuntimeState}; use super::collector::{ FinalizeRequestUpdate, RequestCounterUpdate, StatsAggregatorUpdate, StatsCollectorConfig, @@ -40,12 +39,13 @@ pub(super) struct ModelMetricsState { pub(super) embedding_item_tps_sum: f64, pub(super) max_chat_output_tps: f64, pub(super) max_embedding_item_tps: f64, - pub(super) kv_cache: KvCacheStatsSnapshot, + pub(super) kv_cache: Option, pub(super) input_tps_distribution: TpsDistribution, aggregate_state_counted: bool, pub(super) counter_output_tps_authoritative: bool, pub(super) chunk_usage_stats_observed: bool, pub(super) kv_cache_stats_observed: bool, + pub(super) kv_cache_received_at: Option, pub(super) engine_stream_stats_observed: bool, pub(super) last_stats_event_at: Option, pub(super) stats_observed_at_unix_ms: u64, @@ -57,12 +57,20 @@ pub(super) struct GenerationMetricsState { pub(super) metrics: ModelMetricsState, pub(super) pinned_input_tps: Option, } -#[derive(Debug, Clone, Default, Deserialize, PartialEq, Eq)] +#[derive(Debug, Clone, Default, PartialEq, Eq)] pub(super) struct KvCacheStatsSnapshot { pub(super) model: String, + pub(super) aliases: Vec, pub(super) kv_cache_capacity_tokens: u64, pub(super) kv_cache_used_tokens: u64, pub(super) kv_cache_free_tokens: u64, + pub(super) source_observed_at_unix_ms: u64, + pub(super) complete: bool, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(super) struct KvCacheStatsEnvelope { + pub(super) models: Vec, } struct RequestCounterState { generation: ModelGeneration, @@ -457,6 +465,27 @@ impl StatsAggregator { } } + let kv_cache_ttl = self.config.kv_cache_stats_ttl; + if !kv_cache_ttl.is_zero() { + for (model_id, generation_state) in &mut self.per_model { + let state = &mut generation_state.metrics; + if state.kv_cache_received_at.is_some_and(|received_at| { + now.saturating_duration_since(received_at) >= kv_cache_ttl + }) { + state.kv_cache_received_at = None; + if state.kv_cache.take().is_some() { + state.stats_observed_at_unix_ms = current_unix_millis(); + tracing::warn!( + model_id, + ttl_ms = kv_cache_ttl.as_millis(), + "clearing stale KV cache stats" + ); + push_dirty_model(&mut dirty_models, model_id.clone()); + } + } + } + } + if let Some(metrics) = self.runtime_state.metrics() { metrics .observe_engine_stats_model_states(ENGINE_STATS_SOURCE, self.model_state_count()); @@ -870,6 +899,7 @@ impl ModelMetricsState { pub(super) fn current_stats(&self, inputs: ModelStatsSnapshotInputs) -> CurrentModelStats { let (stats_capabilities, stats_sources) = self.stats_labels(); + let kv_cache = self.kv_cache.as_ref().filter(|snapshot| snapshot.complete); let active_chat_output_tps = if self.counter_output_tps_authoritative { 0.0 } else { @@ -889,9 +919,21 @@ impl ModelMetricsState { max_embedding_item_tps: self.max_embedding_item_tps, queue_size: inputs.queue_size, queued_input_size: inputs.queued_input_size, - kv_cache_capacity_tokens: self.kv_cache.kv_cache_capacity_tokens, - kv_cache_used_tokens: self.kv_cache.kv_cache_used_tokens, - kv_cache_free_tokens: self.kv_cache.kv_cache_free_tokens, + kv_cache_capacity_tokens: kv_cache + .map(|snapshot| snapshot.kv_cache_capacity_tokens) + .unwrap_or_default(), + kv_cache_used_tokens: kv_cache + .map(|snapshot| snapshot.kv_cache_used_tokens) + .unwrap_or_default(), + kv_cache_free_tokens: kv_cache + .map(|snapshot| snapshot.kv_cache_free_tokens) + .unwrap_or_default(), + kv_cache: kv_cache.map(|snapshot| CurrentKvCacheStats { + capacity_tokens: snapshot.kv_cache_capacity_tokens, + used_tokens: snapshot.kv_cache_used_tokens, + free_tokens: snapshot.kv_cache_free_tokens, + source_observed_at_unix_ms: snapshot.source_observed_at_unix_ms, + }), num_running_queries: inputs.num_running_queries, max_engine_concurrency: None, total_query_input_size: inputs.total_query_input_size, diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs index 1c77c6958..39da72677 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs @@ -24,15 +24,17 @@ use crate::runtime_state::ModelGeneration; use crate::{CurrentModelStats, PylonRuntimeState, RequestObservationEvent}; use stargate_runtime::OwnedTask; -use super::aggregator::{ENGINE_STATS_SOURCE, KvCacheStatsSnapshot, StatsAggregator}; +use super::aggregator::{ENGINE_STATS_SOURCE, StatsAggregator}; +use super::kv_stats_stream::{KvStatsStreamConfig, run_kv_stats_stream}; const DEFAULT_OBSERVATION_CHANNEL_CAPACITY: usize = 1024; const DEFAULT_SMOOTHING_WINDOW_SIZE: usize = 8; const DEFAULT_MIN_INPUT_TOKENS: u64 = 1; const DEFAULT_MIN_OUTPUT_TOKENS: u64 = 1; const DEFAULT_DURATION_FLOOR: Duration = Duration::from_millis(10); -const DEFAULT_KV_CACHE_POLL_INTERVAL: Duration = Duration::from_secs(1); -const DEFAULT_KV_CACHE_REQUEST_TIMEOUT: Duration = Duration::from_secs(1); +const DEFAULT_KV_CACHE_RECONNECT_INTERVAL: Duration = Duration::from_secs(1); +const DEFAULT_KV_CACHE_CONNECT_TIMEOUT: Duration = Duration::from_secs(1); +const DEFAULT_KV_CACHE_STATS_TTL: Duration = Duration::from_secs(5); const DEFAULT_ENGINE_STATS_REQUEST_TTL: Duration = Duration::from_secs(300); const DEFAULT_ENGINE_STATS_MODEL_TTL: Duration = Duration::from_secs(30); const DEFAULT_ENGINE_STATS_SWEEP_INTERVAL: Duration = Duration::from_secs(1); @@ -45,8 +47,9 @@ pub struct StatsCollectorConfig { pub min_output_tokens: u64, pub duration_floor: Duration, pub kv_cache_stats_url: Option, - pub kv_cache_poll_interval: Duration, - pub kv_cache_request_timeout: Duration, + pub kv_cache_reconnect_interval: Duration, + pub kv_cache_connect_timeout: Duration, + pub kv_cache_stats_ttl: Duration, pub engine_stats_request_ttl: Duration, pub engine_stats_model_ttl: Duration, pub engine_stats_sweep_interval: Duration, @@ -62,8 +65,9 @@ impl Default for StatsCollectorConfig { min_output_tokens: DEFAULT_MIN_OUTPUT_TOKENS, duration_floor: DEFAULT_DURATION_FLOOR, kv_cache_stats_url: None, - kv_cache_poll_interval: DEFAULT_KV_CACHE_POLL_INTERVAL, - kv_cache_request_timeout: DEFAULT_KV_CACHE_REQUEST_TIMEOUT, + kv_cache_reconnect_interval: DEFAULT_KV_CACHE_RECONNECT_INTERVAL, + kv_cache_connect_timeout: DEFAULT_KV_CACHE_CONNECT_TIMEOUT, + kv_cache_stats_ttl: DEFAULT_KV_CACHE_STATS_TTL, engine_stats_request_ttl: DEFAULT_ENGINE_STATS_REQUEST_TTL, engine_stats_model_ttl: DEFAULT_ENGINE_STATS_MODEL_TTL, engine_stats_sweep_interval: DEFAULT_ENGINE_STATS_SWEEP_INTERVAL, @@ -329,9 +333,31 @@ async fn run_stats_collector( ) { let config = aggregator.config.clone(); let runtime_state = aggregator.runtime_state.clone(); - let http_client = reqwest::Client::new(); - let mut kv_cache_poll = tokio::time::interval(config.kv_cache_poll_interval); - kv_cache_poll.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let kv_stream_stop = stop.child_token(); + let (mut kv_stats_rx, kv_stream_task) = if let Some(url) = config.kv_cache_stats_url.clone() { + let (tx, rx) = flume::bounded(config.observation_channel_capacity); + let stream_config = KvStatsStreamConfig { + url, + reconnect_interval: config.kv_cache_reconnect_interval, + connect_timeout: config.kv_cache_connect_timeout, + idle_timeout: if config.kv_cache_stats_ttl.is_zero() { + DEFAULT_KV_CACHE_STATS_TTL + } else { + config.kv_cache_stats_ttl + }, + }; + let stream_stop = kv_stream_stop.clone(); + ( + Some(rx), + Some(tokio::spawn(run_kv_stats_stream( + stream_config, + tx, + stream_stop, + ))), + ) + } else { + (None, None) + }; let mut engine_stats_sweep = tokio::time::interval(config.engine_stats_sweep_interval); engine_stats_sweep.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); let mut stats_aggregator_updated_models = Vec::with_capacity(2); @@ -364,6 +390,19 @@ async fn run_stats_collector( } apply_collector_command(&mut aggregator, &runtime_state, command); } + snapshot = async { + match &kv_stats_rx { + Some(rx) => rx.recv_async().await.ok(), + None => std::future::pending().await, + } + } => { + let Some(snapshot) = snapshot else { + kv_stats_rx = None; + continue; + }; + let updates = aggregator.apply_kv_cache_snapshot(snapshot, TokioInstant::now()); + publish_model_stats_updates(&runtime_state, updates); + } event = observation_rx.recv_async() => { let Ok(event) = event else { break 'collector; @@ -418,34 +457,13 @@ async fn run_stats_collector( } publish_model_stats_updates(&runtime_state, updated_models); } - _ = kv_cache_poll.tick(), if config.kv_cache_stats_url.is_some() => { - let Some(kv_cache) = stop - .run_until_cancelled(poll_kv_cache_stats(&config, &http_client)) - .await - else { - break 'collector; - }; - let Some(kv_cache) = kv_cache else { - continue; - }; - if kv_cache.model.is_empty() { - tracing::warn!("dropping KV-cache stats without model id"); - continue; - } - let model_id = kv_cache.model.clone(); - let Some((model_id, updated_stats)) = aggregator.apply_kv_cache_stats(kv_cache) - else { - tracing::warn!( - model_id, - configured_models = ?aggregator.per_model.keys(), - "dropping KV-cache stats for a model with no live generation" - ); - continue; - }; - publish_model_stats_update(&runtime_state, model_id, updated_stats); - } } } + + kv_stream_stop.cancel(); + if let Some(task) = kv_stream_task { + let _ = task.await; + } } fn publish_observation_event( @@ -516,33 +534,6 @@ fn observe_aggregate_counts(aggregator: &StatsAggregator, runtime_state: &PylonR } } -async fn poll_kv_cache_stats( - config: &StatsCollectorConfig, - http_client: &reqwest::Client, -) -> Option { - let url = config.kv_cache_stats_url.as_ref()?; - let response = http_client - .get(url) - .timeout(config.kv_cache_request_timeout) - .send() - .await - .inspect_err(|error| { - tracing::warn!(url, error = %error, "failed to poll KV-cache stats"); - }) - .ok()?; - if !response.status().is_success() { - tracing::warn!(url, status = %response.status(), "KV-cache stats endpoint returned non-success status"); - return None; - } - response - .json() - .await - .inspect_err(|error| { - tracing::warn!(url, error = %error, "failed to parse KV-cache stats"); - }) - .ok() -} - fn drain_ready(rx: &flume::Receiver, mut consume: impl FnMut(T)) { for _ in 0..rx.len() { let Ok(value) = rx.try_recv() else { break }; @@ -564,8 +555,9 @@ fn retain_latest_model_updates( mod tests { use std::collections::HashMap; use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; - use super::super::aggregator::{KvCacheStatsSnapshot, StatsAggregator}; + use super::super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot, StatsAggregator}; use super::super::metrics::PylonMetrics; use super::super::projection::fallback_update_from_observation; use super::*; @@ -1024,9 +1016,18 @@ mod tests { fn kv_cache_stats(model: &str) -> KvCacheStatsSnapshot { KvCacheStatsSnapshot { model: model.to_string(), + aliases: Vec::new(), kv_cache_capacity_tokens: 1_000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600, + source_observed_at_unix_ms: 1, + complete: true, + } + } + + fn kv_cache_envelope(model: KvCacheStatsSnapshot) -> KvCacheStatsEnvelope { + KvCacheStatsEnvelope { + models: vec![model], } } @@ -2326,7 +2327,7 @@ mod tests { ); #[test] - fn snapshot_includes_polled_kv_cache_stats() { + fn snapshot_includes_streamed_kv_cache_stats() { let mut aggregator = test_aggregator(StatsCollectorConfig::default()); let stats = aggregator .apply_kv_cache_stats(kv_cache_stats("model-a")) @@ -2335,23 +2336,63 @@ mod tests { assert_stats!(stats; kv_cache_capacity_tokens: 1_000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600); } + #[test] + fn kv_cache_snapshot_matches_alias_and_replaces_atomically() { + let mut aggregator = test_aggregator(StatsCollectorConfig::default()); + let mut snapshot = kv_cache_stats("canonical-model"); + snapshot.aliases.push("model-a".to_string()); + let updates = + aggregator.apply_kv_cache_snapshot(kv_cache_envelope(snapshot), TokioInstant::now()); + let stats = published_stats(updates); + assert_stats!(stats; kv_cache_capacity_tokens: 1_000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600); + assert!(stats.kv_cache.is_some()); + + let mut incomplete = kv_cache_stats("canonical-model"); + incomplete.aliases.push("model-a".to_string()); + incomplete.complete = false; + let updates = + aggregator.apply_kv_cache_snapshot(kv_cache_envelope(incomplete), TokioInstant::now()); + let stats = published_stats(updates); + assert_stats!(stats; kv_cache_capacity_tokens: 0, kv_cache_used_tokens: 0, kv_cache_free_tokens: 0); + assert!(stats.kv_cache.is_none()); + } + + #[test] + fn stale_kv_cache_snapshot_expires_without_clearing_request_stats() { + let config = config!(kv_cache_stats_ttl: milliseconds(10)); + let mut aggregator = test_aggregator(config); + aggregator.stream("req", (0, 0), false, Duration::ZERO); + let request_updates = aggregator.stream("req", (10, 4), false, milliseconds(100)); + assert_eq!(published_stats(request_updates).output_tps, 40.0); + + let received_at = TokioInstant::now(); + aggregator + .apply_kv_cache_snapshot(kv_cache_envelope(kv_cache_stats("model-a")), received_at); + let updates = aggregator.sweep_stale(received_at + milliseconds(11)); + let stats = published_stats(updates); + assert!(stats.kv_cache.is_none()); + assert_eq!(stats.kv_cache_capacity_tokens, 0); + assert_eq!(stats.output_tps, 40.0); + } + #[tokio::test] - async fn kv_cache_poll_updates_model_metrics() { - async fn kv_cache_stats() -> Json { - Json(serde_json::json!({ - "model": "model-a", - "kv_cache_capacity_tokens": 1000, - "kv_cache_used_tokens": 400, - "kv_cache_free_tokens": 600 - })) - } + async fn kv_cache_stream_updates_model_metrics() { + const SNAPSHOT: &str = "{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":42,\"models\":[{\"model\":\"model-a\",\"aliases\":[],\"routing_cache\":{\"role\":\"decode\",\"capacity_tokens\":1000,\"used_tokens\":400,\"free_tokens\":600,\"complete\":true},\"pools\":[]}]}\n"; + let requests = Arc::new(AtomicUsize::new(0)); + let handler_requests = requests.clone(); let metrics = PylonMetrics::new().expect("metrics should initialize"); - let app = Router::new().route("/kv-cache", get(kv_cache_stats)); + let app = Router::new().route( + "/kv-cache", + get(move || { + handler_requests.fetch_add(1, Ordering::Relaxed); + async { SNAPSHOT } + }), + ); let (addr, server) = spawn_kv_cache_server(app).await; let config = config!( kv_cache_stats_url: Some(format!("http://{addr}/kv-cache")), - kv_cache_poll_interval: milliseconds(10), - kv_cache_request_timeout: seconds(1), + kv_cache_reconnect_interval: milliseconds(10), + kv_cache_connect_timeout: seconds(1), ); let collector = RunningCollector::spawn(config, Some(metrics.clone()), false); let stats = collector @@ -2360,10 +2401,24 @@ mod tests { }) .await; assert_stats!(stats; kv_cache_capacity_tokens: 1000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600); + assert_eq!( + stats + .kv_cache + .as_ref() + .map(|stats| stats.source_observed_at_unix_ms), + Some(42) + ); let body = metrics.gather_text().expect("metrics should encode"); assert!(body.contains(r#"pylon_model_kv_cache_capacity_tokens{model="model-a"} 1000"#)); assert!(body.contains(r#"pylon_model_kv_cache_used_tokens{model="model-a"} 400"#)); assert!(body.contains(r#"pylon_model_kv_cache_free_tokens{model="model-a"} 600"#)); + tokio::time::timeout(seconds(1), async { + while requests.load(Ordering::Relaxed) < 2 { + tokio::time::sleep(milliseconds(10)).await; + } + }) + .await + .expect("finite KV stats responses should reconnect"); tokio::time::timeout(seconds(2), collector.handle.shutdown()) .await .expect("collector should stop"); @@ -2371,7 +2426,7 @@ mod tests { } #[tokio::test] - async fn stats_collector_shutdown_interrupts_blocked_kv_cache_poll() { + async fn stats_collector_shutdown_interrupts_blocked_kv_cache_connect() { let poll_entered = Arc::new(tokio::sync::Barrier::new(2)); let server_poll_entered = poll_entered.clone(); let app = Router::new().route( @@ -2387,14 +2442,14 @@ mod tests { let (addr, server) = spawn_kv_cache_server(app).await; let config = config!( kv_cache_stats_url: Some(format!("http://{addr}/kv-cache")), - kv_cache_poll_interval: milliseconds(1), - kv_cache_request_timeout: seconds(60), + kv_cache_reconnect_interval: milliseconds(1), + kv_cache_connect_timeout: seconds(60), ); let collector = RunningCollector::spawn(config, None, false); poll_entered.wait().await; let stopped = tokio::time::timeout(seconds(1), collector.handle.shutdown()).await; server.abort(); - stopped.expect("collector shutdown should interrupt blocked KV-cache poll"); + stopped.expect("collector shutdown should interrupt blocked KV-cache connect"); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs new file mode 100644 index 000000000..3641cb274 --- /dev/null +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs @@ -0,0 +1,379 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::collections::HashSet; +use std::time::Duration; + +use bytes::Bytes; +use futures::{Stream, StreamExt}; +use serde::Deserialize; +use tokio_util::sync::CancellationToken; + +use super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot}; + +const MAX_LINE_BYTES: usize = 1024 * 1024; + +pub(super) struct KvStatsStreamConfig { + pub(super) url: String, + pub(super) reconnect_interval: Duration, + pub(super) connect_timeout: Duration, + pub(super) idle_timeout: Duration, +} + +#[derive(Deserialize)] +struct RawSnapshot { + v: u8, + #[serde(rename = "type")] + event_type: String, + observed_at_unix_ms: u64, + models: Vec, +} + +#[derive(Deserialize)] +struct RawModelStats { + model: String, + #[serde(default)] + aliases: Vec, + routing_cache: Option, +} + +#[derive(Deserialize)] +struct RawRoutingCacheStats { + capacity_tokens: Option, + used_tokens: Option, + free_tokens: Option, + complete: bool, +} + +pub(super) fn parse_kv_stats_snapshot(line: &[u8]) -> anyhow::Result { + let snapshot: RawSnapshot = serde_json::from_slice(line)?; + anyhow::ensure!( + snapshot.v == 1, + "unsupported KV stats version {}", + snapshot.v + ); + anyhow::ensure!( + snapshot.event_type == "kv_stats_snapshot", + "unsupported KV stats event type {}", + snapshot.event_type + ); + let mut identities = HashSet::new(); + let models = snapshot + .models + .into_iter() + .map(|model| -> anyhow::Result<_> { + anyhow::ensure!(!model.model.trim().is_empty(), "KV stats model is empty"); + anyhow::ensure!( + identities.insert(model.model.clone()), + "duplicate KV stats identity {}", + model.model + ); + for alias in &model.aliases { + anyhow::ensure!(!alias.trim().is_empty(), "KV stats alias is empty"); + anyhow::ensure!( + identities.insert(alias.clone()), + "duplicate KV stats identity {alias}" + ); + } + let routing = model.routing_cache; + let complete = snapshot.observed_at_unix_ms > 0 + && routing.as_ref().is_some_and(|routing| { + let Some((capacity, used, free)) = routing + .capacity_tokens + .zip(routing.used_tokens) + .zip(routing.free_tokens) + .map(|((capacity, used), free)| (capacity, used, free)) + else { + return false; + }; + routing.complete && capacity > 0 && used.checked_add(free) == Some(capacity) + }); + let (capacity, used, free) = routing + .map(|routing| { + ( + routing.capacity_tokens.unwrap_or_default(), + routing.used_tokens.unwrap_or_default(), + routing.free_tokens.unwrap_or_default(), + ) + }) + .unwrap_or_default(); + Ok(KvCacheStatsSnapshot { + model: model.model, + aliases: model.aliases, + kv_cache_capacity_tokens: capacity, + kv_cache_used_tokens: used, + kv_cache_free_tokens: free, + source_observed_at_unix_ms: snapshot.observed_at_unix_ms, + complete, + }) + }) + .collect::>>()?; + Ok(KvCacheStatsEnvelope { models }) +} + +pub(super) async fn run_kv_stats_stream( + config: KvStatsStreamConfig, + updates: flume::Sender, + stop: CancellationToken, +) { + let client = reqwest::Client::new(); + while !stop.is_cancelled() { + if let Err(error) = read_stream_once(&config, &client, &updates, &stop).await { + tracing::warn!(url = config.url, %error, "KV stats stream disconnected"); + } + if stop + .run_until_cancelled(tokio::time::sleep(config.reconnect_interval)) + .await + .is_none() + { + break; + } + } +} + +async fn read_stream_once( + config: &KvStatsStreamConfig, + client: &reqwest::Client, + updates: &flume::Sender, + stop: &CancellationToken, +) -> anyhow::Result<()> { + let response = tokio::select! { + _ = stop.cancelled() => return Ok(()), + response = tokio::time::timeout( + config.connect_timeout, + client + .get(&config.url) + .header(reqwest::header::ACCEPT, "application/x-ndjson") + .send(), + ) => response??, + }; + anyhow::ensure!( + response.status().is_success(), + "KV stats endpoint returned {}", + response.status() + ); + drain_response(response.bytes_stream(), updates, stop, config.idle_timeout).await +} + +async fn drain_response( + mut stream: S, + updates: &flume::Sender, + stop: &CancellationToken, + idle_timeout: Duration, +) -> anyhow::Result<()> +where + S: Stream> + Unpin, +{ + let mut buffer = Vec::with_capacity(4096); + let mut discarding_oversized_line = false; + loop { + let chunk = tokio::select! { + _ = stop.cancelled() => return Ok(()), + chunk = tokio::time::timeout(idle_timeout, stream.next()) => { + chunk.map_err(|_| anyhow::anyhow!("KV stats stream became idle"))? + }, + }; + let Some(chunk) = chunk else { + anyhow::bail!("KV stats stream ended"); + }; + let chunk = chunk?; + let mut remaining = chunk.as_ref(); + while let Some(newline) = remaining.iter().position(|byte| *byte == b'\n') { + let segment = &remaining[..newline]; + remaining = &remaining[newline + 1..]; + if discarding_oversized_line { + discarding_oversized_line = false; + continue; + } + if buffer.len().saturating_add(segment.len()) > MAX_LINE_BYTES { + tracing::warn!("dropping oversized KV stats line"); + buffer.clear(); + continue; + } + buffer.extend_from_slice(segment); + if buffer.iter().all(u8::is_ascii_whitespace) { + buffer.clear(); + continue; + } + match parse_kv_stats_snapshot(&buffer) { + Ok(snapshot) => { + match stop.run_until_cancelled(updates.send_async(snapshot)).await { + None | Some(Err(_)) => return Ok(()), + Some(Ok(())) => {} + } + } + Err(error) => tracing::warn!(%error, "dropping invalid KV stats snapshot"), + } + buffer.clear(); + } + if discarding_oversized_line { + continue; + } + if buffer.len().saturating_add(remaining.len()) > MAX_LINE_BYTES { + tracing::warn!("dropping oversized KV stats line"); + buffer.clear(); + discarding_oversized_line = true; + } else { + buffer.extend_from_slice(remaining); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_complete_snapshot_without_placement_data() { + let snapshot = parse_kv_stats_snapshot( + br#"{"v":1,"type":"kv_stats_snapshot","snapshot_id":9,"observed_at_unix_ms":42,"models":[{"model":"m","aliases":["alias"],"routing_cache":{"role":"decode","capacity_tokens":100,"used_tokens":40,"free_tokens":60,"complete":true},"pools":[]}]}"#, + ) + .unwrap(); + assert_eq!(snapshot.models[0].source_observed_at_unix_ms, 42); + assert_eq!(snapshot.models.len(), 1); + assert_eq!(snapshot.models[0].aliases, ["alias"]); + assert!(snapshot.models[0].complete); + } + + #[test] + fn inconsistent_complete_snapshot_is_not_usable() { + let snapshot = parse_kv_stats_snapshot( + br#"{"v":1,"type":"kv_stats_snapshot","observed_at_unix_ms":42,"models":[{"model":"m","aliases":[],"routing_cache":{"capacity_tokens":100,"used_tokens":80,"free_tokens":30,"complete":true},"pools":[]}]}"#, + ) + .unwrap(); + assert!(!snapshot.models[0].complete); + } + + #[test] + fn duplicate_alias_ownership_rejects_the_whole_snapshot() { + let result = parse_kv_stats_snapshot( + br#"{"v":1,"type":"kv_stats_snapshot","observed_at_unix_ms":42,"models":[{"model":"a","aliases":["shared"],"routing_cache":null},{"model":"b","aliases":["shared"],"routing_cache":null}]}"#, + ); + assert!(result.is_err()); + } + + #[test] + fn line_limit_accepts_a_representative_thousand_model_snapshot() { + let models = (0..1_000) + .map(|index| { + serde_json::json!({ + "model": format!("model-{index:04}"), + "aliases": [format!("deployment-model-{index:04}")], + "routing_cache": { + "role": "decode", + "capacity_tokens": 65_536_000, + "used_tokens": 6_553_600, + "free_tokens": 58_982_400, + "complete": true + }, + "pools": [{ + "namespace": "dynamo", + "component": "backend", + "endpoint": "generate", + "role": "decode", + "storage_tier": "device", + "block_size_tokens": 64, + "expected_ranks": 8, + "observed_ranks": 8, + "capacity_blocks": 1_024_000, + "used_blocks": 102_400, + "free_blocks": 921_600, + "active_decode_blocks": 81_920, + "complete": true + }] + }) + }) + .collect::>(); + let line = serde_json::to_vec(&serde_json::json!({ + "v": 1, + "type": "kv_stats_snapshot", + "snapshot_id": 1, + "observed_at_unix_ms": 1, + "models": models + })) + .unwrap(); + + assert!( + line.len() <= MAX_LINE_BYTES, + "representative snapshot is {} bytes", + line.len() + ); + } + + #[tokio::test] + async fn fragmented_ndjson_is_reassembled_before_publication() { + let chunks = futures::stream::iter([ + Ok::<_, reqwest::Error>(Bytes::from_static( + b"{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":9,", + )), + Ok(Bytes::from_static( + b"\"models\":[{\"model\":\"m\",\"routing_cache\":{\"capacity_tokens\":10,\"used_tokens\":4,\"free_tokens\":6,\"complete\":true}}]}\n", + )), + ]); + let (tx, rx) = flume::bounded(1); + let result = drain_response( + chunks, + &tx, + &CancellationToken::new(), + Duration::from_secs(1), + ) + .await; + assert!( + result.is_err(), + "finite response should end after publishing" + ); + let snapshot = rx.try_recv().expect("snapshot should be published"); + assert_eq!(snapshot.models[0].source_observed_at_unix_ms, 9); + assert!(snapshot.models[0].complete); + } + + #[tokio::test] + async fn oversized_line_is_dropped_without_buffering_or_losing_the_next_snapshot() { + let valid = + b"{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":9,\"models\":[]}\n"; + let mut first = vec![b'x'; MAX_LINE_BYTES + 1]; + first.extend_from_slice(b"\n"); + let chunks = futures::stream::iter([ + Ok::<_, reqwest::Error>(Bytes::from(first)), + Ok(Bytes::from_static(valid)), + ]); + let (tx, rx) = flume::bounded(1); + + let result = drain_response( + chunks, + &tx, + &CancellationToken::new(), + Duration::from_secs(1), + ) + .await; + + assert!( + result.is_err(), + "finite response should end after publishing" + ); + assert_eq!( + rx.try_recv() + .expect("valid line should still publish") + .models + .len(), + 0 + ); + } + + #[tokio::test] + async fn idle_stream_is_reconnected() { + let stream = futures::stream::pending::>(); + let (tx, _rx) = flume::bounded(1); + + let error = drain_response( + stream, + &tx, + &CancellationToken::new(), + Duration::from_millis(10), + ) + .await + .unwrap_err(); + + assert!(error.to_string().contains("became idle")); + } +} diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs index bdcdd9725..020e4471b 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs @@ -36,6 +36,7 @@ macro_rules! owned_task_handle { mod aggregator; mod collector; mod engine_stats_stream; +mod kv_stats_stream; mod metrics; mod projection; pub(crate) mod token_metrics; diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs index d416c86e5..a12bd81d5 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs @@ -21,9 +21,9 @@ use crate::runtime_state::ModelGeneration; use crate::{CurrentModelStats, RequestObservation, RequestObservationEvent}; use super::aggregator::{ - EmbeddingThroughputSample, InputThroughputSample, KvCacheStatsSnapshot, ModelMetricsState, - ModelStatsSnapshotInputs, StatsAggregator, apply_input_throughput_sample, current_unix_millis, - output_decode_duration, push_sample, tps_for_units, + EmbeddingThroughputSample, InputThroughputSample, KvCacheStatsEnvelope, KvCacheStatsSnapshot, + ModelMetricsState, ModelStatsSnapshotInputs, StatsAggregator, apply_input_throughput_sample, + current_unix_millis, output_decode_duration, push_sample, tps_for_units, }; use super::collector::{ FinalizeRequestUpdate, RequestCounterUpdate, StatsAggregatorUpdate, StatsCollectorConfig, @@ -92,20 +92,52 @@ impl StatsAggregator { }) } + #[cfg(test)] pub(super) fn apply_kv_cache_stats( &mut self, kv_cache: KvCacheStatsSnapshot, ) -> Option { let model_id = kv_cache.model.clone(); let model_state = self.per_model.get_mut(&model_id)?; - model_state.metrics.kv_cache = kv_cache; + model_state.metrics.kv_cache = valid_kv_cache(&kv_cache).then_some(kv_cache); model_state.metrics.kv_cache_stats_observed = true; + model_state.metrics.kv_cache_received_at = Some(TokioInstant::now()); model_state.metrics.stats_observed_at_unix_ms = current_unix_millis(); let generation = model_state.generation.clone(); let stats = self.snapshot(&model_id); Some((generation, stats)) } + pub(super) fn apply_kv_cache_snapshot( + &mut self, + snapshot: KvCacheStatsEnvelope, + received_at: TokioInstant, + ) -> Vec { + let model_ids = self.per_model.keys().cloned().collect::>(); + let mut changed = Vec::new(); + for model_id in model_ids { + let observed = snapshot.models.iter().find(|model| { + model.model == model_id || model.aliases.iter().any(|alias| alias == &model_id) + }); + let next = observed.filter(|model| valid_kv_cache(model)).cloned(); + let model_state = self + .per_model + .get_mut(&model_id) + .expect("model id came from the current generation map"); + if model_state.metrics.kv_cache == next { + model_state.metrics.kv_cache_received_at = next.as_ref().map(|_| received_at); + continue; + } + model_state.metrics.kv_cache = next; + model_state.metrics.kv_cache_received_at = + model_state.metrics.kv_cache.as_ref().map(|_| received_at); + model_state.metrics.kv_cache_stats_observed |= observed.is_some(); + model_state.metrics.stats_observed_at_unix_ms = current_unix_millis(); + changed.push(model_id); + } + self.snapshots(changed) + } + pub(super) fn snapshot(&self, model_id: &str) -> CurrentModelStats { let queue = self.runtime_state.snapshot_live_model(model_id); let inputs = ModelStatsSnapshotInputs { @@ -253,6 +285,15 @@ impl StatsAggregator { } } +fn valid_kv_cache(snapshot: &KvCacheStatsSnapshot) -> bool { + snapshot.complete + && snapshot.kv_cache_capacity_tokens > 0 + && snapshot + .kv_cache_used_tokens + .checked_add(snapshot.kv_cache_free_tokens) + == Some(snapshot.kv_cache_capacity_tokens) +} + pub(super) fn fallback_update_from_observation( observation: &RequestObservation, generation: Option, diff --git a/src/libraries/rust/stargate/crates/pylon/src/main.rs b/src/libraries/rust/stargate/crates/pylon/src/main.rs index f900b7310..ff64f84fa 100644 --- a/src/libraries/rust/stargate/crates/pylon/src/main.rs +++ b/src/libraries/rust/stargate/crates/pylon/src/main.rs @@ -110,8 +110,8 @@ struct Args { /// Timeout for calibration requests in milliseconds #[arg(long, default_value_t = 30000, value_name = "MS")] bringup_calibration_timeout_ms: u64, - /// Upstream HTTP path to poll for KV-cache stats. Omit to disable KV metric polling - #[arg(long, value_name = "PATH")] + /// Upstream HTTP path for the canonical KV-cache stats stream + #[arg(long, default_value = "/v1/kv-cache/stats/stream", value_name = "PATH")] kv_cache_stats_path: Option, /// Engine stats stream source selection mode #[arg(long, default_value_t = EngineStatsStreamMode::Auto, value_name = "MODE")] @@ -587,7 +587,10 @@ mod tests { assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Auto); assert_eq!(args.engine_stats_stream_path, "/pylon/v1/stats/stream"); - assert!(metrics_config.kv_cache_stats_url.is_none()); + assert_eq!( + metrics_config.kv_cache_stats_url.as_deref(), + Some("http://127.0.0.1:8090/v1/kv-cache/stats/stream") + ); assert!( !metrics_config.openai_fallback_stats_enabled, "auto mode should wait for a permanent unsupported stream response before fallback stats" @@ -601,12 +604,15 @@ mod tests { let metrics_config = stats_collector_config_from_args(&args, &upstream); assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Off); - assert!(metrics_config.kv_cache_stats_url.is_none()); + assert_eq!( + metrics_config.kv_cache_stats_url.as_deref(), + Some("http://127.0.0.1:8090/v1/kv-cache/stats/stream") + ); assert!(metrics_config.openai_fallback_stats_enabled); } #[test] - fn kv_cache_stats_path_enables_explicit_kv_cache_polling() { + fn kv_cache_stats_path_overrides_the_canonical_stream() { let args = parse_args("--kv-cache-stats-path /kv-cache/stats"); let upstream = normalize_base_url(&args.upstream_http_base_url); let metrics_config = stats_collector_config_from_args(&args, &upstream); diff --git a/src/libraries/rust/stargate/crates/pylon/src/startup.rs b/src/libraries/rust/stargate/crates/pylon/src/startup.rs index 00c3d02fd..6d99fe1fb 100644 --- a/src/libraries/rust/stargate/crates/pylon/src/startup.rs +++ b/src/libraries/rust/stargate/crates/pylon/src/startup.rs @@ -601,8 +601,7 @@ pub(crate) fn stats_collector_config_from_args( ) -> StatsCollectorConfig { StatsCollectorConfig { openai_fallback_stats_enabled: args.engine_stats_stream == EngineStatsStreamMode::Off, - // Mock benchmark backends can expose live KV-cache occupancy over HTTP; - // real upstreams usually do not, so polling is explicit. + // Dynamo exposes this canonical stream independently from request stats. kv_cache_stats_url: args.kv_cache_stats_path.as_deref().map(|path| { format!( "{}/{}", diff --git a/src/libraries/rust/stargate/crates/stargate/src/metrics.rs b/src/libraries/rust/stargate/crates/stargate/src/metrics.rs index 5e2f3d514..eb0d89a6e 100644 --- a/src/libraries/rust/stargate/crates/stargate/src/metrics.rs +++ b/src/libraries/rust/stargate/crates/stargate/src/metrics.rs @@ -114,6 +114,7 @@ define_stargate_metrics! { gauges { active_inference_servers("active_inference_servers", "Active inference servers available for a routing target", ["routing_key", "model"]); tls_certificate_expiry_seconds("tls_certificate_expiry_seconds", "Unix timestamp when the active TLS certificate expires", ["material_type"]); + kv_cache_stats_disagreement("kv_cache_stats_disagreement", "Whether complete KV cache observations disagree for a routing target", ["routing_key", "model"]); } } @@ -209,6 +210,18 @@ impl StargateMetrics { .with_label_values(&[routing_key.unwrap_or(""), model]) .set(count.try_into().unwrap_or(i64::MAX)); } + + #[inline] + pub fn set_kv_cache_stats_disagreement( + &self, + routing_key: Option<&str>, + model: &str, + disagrees: bool, + ) { + self.kv_cache_stats_disagreement + .with_label_values(&[routing_key.unwrap_or(""), model]) + .set(i64::from(disagrees)); + } } // -- Metrics HTTP server ----------------------------------------------------- @@ -290,6 +303,7 @@ mod tests { "pulsar-wait-and-widen", ) .inc(); + metrics.set_kv_cache_stats_disagreement(Some("routing-a"), "model-a", true); let body = metrics.gather_text().expect("metrics should encode"); assert!( @@ -312,6 +326,12 @@ mod tests { body.contains("llm_request_router_routing_kv_free_token_fallback_selections_total"), "custom KV-free-token fallback counter prefix missing:\n{body}" ); + assert!( + body.contains( + r#"llm_request_router_kv_cache_stats_disagreement{model="model-a",routing_key="routing-a"} 1"# + ), + "KV-cache disagreement gauge missing:\n{body}" + ); assert!( !body.contains("stargate_requests_total"), "default stargate prefix leaked into custom metric output:\n{body}" diff --git a/src/libraries/rust/stargate/crates/stargate/src/routing_state/cluster_snapshots.rs b/src/libraries/rust/stargate/crates/stargate/src/routing_state/cluster_snapshots.rs index ed781cbe0..16dade548 100644 --- a/src/libraries/rust/stargate/crates/stargate/src/routing_state/cluster_snapshots.rs +++ b/src/libraries/rust/stargate/crates/stargate/src/routing_state/cluster_snapshots.rs @@ -33,15 +33,85 @@ use super::snapshots::{ fn set_cluster_scoped_stats(stats: &mut ModelStats, src: &ModelStats) { stats.max_output_tps = src.max_output_tps; - stats.kv_cache_capacity_tokens = src.kv_cache_capacity_tokens; - stats.kv_cache_used_tokens = src.kv_cache_used_tokens; - stats.kv_cache_free_tokens = src.kv_cache_free_tokens; stats.num_running_queries = src.num_running_queries; stats.max_engine_concurrency = src.max_engine_concurrency; stats.total_query_input_size = src.total_query_input_size; stats.queue_time_estimate_ms_by_priority = src.queue_time_estimate_ms_by_priority.clone(); } +fn valid_kv_cache(stats: &ModelStats) -> bool { + stats.kv_cache.as_ref().is_some_and(|kv_cache| { + kv_cache.complete + && kv_cache.source_observed_at_unix_ms > 0 + && kv_cache.capacity_tokens > 0 + && kv_cache.used_tokens.checked_add(kv_cache.free_tokens) + == Some(kv_cache.capacity_tokens) + }) +} + +fn kv_cache_stats_disagree(backends: &[Arc]) -> bool { + let mut values = backends + .iter() + .filter(|backend| valid_kv_cache(&backend.stats)) + .map(|backend| { + let stats = backend.stats.kv_cache.as_ref().unwrap(); + (stats.capacity_tokens, stats.used_tokens, stats.free_tokens) + }); + let Some(first) = values.next() else { + return false; + }; + values.any(|value| value != first) +} + +fn set_cluster_kv_stats( + stats: &mut ModelStats, + backends: &[Arc], + legacy_source: &ModelStats, +) { + let selected = backends + .iter() + .filter(|backend| valid_kv_cache(&backend.stats)) + .max_by(|left, right| { + let left_time = left + .stats + .kv_cache + .as_ref() + .map(|stats| stats.source_observed_at_unix_ms) + .unwrap_or_default(); + let right_time = right + .stats + .kv_cache + .as_ref() + .map(|stats| stats.source_observed_at_unix_ms) + .unwrap_or_default(); + (left_time, left.inference_server_id.as_str()) + .cmp(&(right_time, right.inference_server_id.as_str())) + }); + if let Some(selected) = selected { + let kv_cache = selected + .stats + .kv_cache + .expect("selected backend has valid KV stats"); + stats.kv_cache_capacity_tokens = kv_cache.capacity_tokens; + stats.kv_cache_used_tokens = kv_cache.used_tokens; + stats.kv_cache_free_tokens = kv_cache.free_tokens; + stats.kv_cache = Some(kv_cache); + } else if backends + .iter() + .all(|backend| backend.stats.kv_cache.is_none()) + { + stats.kv_cache_capacity_tokens = legacy_source.kv_cache_capacity_tokens; + stats.kv_cache_used_tokens = legacy_source.kv_cache_used_tokens; + stats.kv_cache_free_tokens = legacy_source.kv_cache_free_tokens; + stats.kv_cache = None; + } else { + stats.kv_cache_capacity_tokens = 0; + stats.kv_cache_used_tokens = 0; + stats.kv_cache_free_tokens = 0; + stats.kv_cache = None; + } +} + fn append_unique_strings(target: &mut Vec, values: &[String]) { for value in values { if !target.contains(value) { @@ -187,6 +257,7 @@ impl ClusterRoutingGeneration { } let mut stats = backend_stats; set_cluster_scoped_stats(&mut stats, &source_backend.stats); + set_cluster_kv_stats(&mut stats, &self.backends, &source_backend.stats); let base_snapshot = RoutedClusterSnapshot { cluster_id: source_backend.cluster_id.clone(), stats, @@ -247,6 +318,10 @@ impl RoutedClusterState { } } + pub(super) fn kv_cache_stats_disagree(&self) -> bool { + kv_cache_stats_disagree(&self.generation.lock().backends) + } + pub(super) fn upsert_backend( &self, backend: Arc, diff --git a/src/libraries/rust/stargate/crates/stargate/src/routing_state/clusters.rs b/src/libraries/rust/stargate/crates/stargate/src/routing_state/clusters.rs index fde99b08d..58a73f587 100644 --- a/src/libraries/rust/stargate/crates/stargate/src/routing_state/clusters.rs +++ b/src/libraries/rust/stargate/crates/stargate/src/routing_state/clusters.rs @@ -141,6 +141,15 @@ impl RoutingTargetState { } } + fn kv_cache_stats_disagree(&self) -> bool { + match &*self.generation.lock() { + RoutingTargetGeneration::Active { clusters, .. } => clusters + .values() + .any(|cluster| cluster.kv_cache_stats_disagree()), + RoutingTargetGeneration::Retired => false, + } + } + fn retire_if_empty(&self) -> bool { let mut generation = self.generation.lock(); let RoutingTargetGeneration::Active { @@ -245,25 +254,35 @@ impl RoutingLifecycle { targets } - async fn publish_active_backend_count(&self, target: &RoutingTargetKey) { + async fn publish_target_metrics(&self, target: &RoutingTargetKey) { let Some(metrics) = &self.metrics else { return; }; loop { let before = self.target_state(target).await; - let count = before - .as_ref() - .map_or(0, |target_state| target_state.active_backend_count()); + let (count, kv_cache_disagrees) = before.as_ref().map_or((0, false), |target_state| { + ( + target_state.active_backend_count(), + target_state.kv_cache_stats_disagree(), + ) + }); metrics.set_active_inference_servers( target.routing_key.as_deref(), &target.model_id, count, ); + metrics.set_kv_cache_stats_disagreement( + target.routing_key.as_deref(), + &target.model_id, + kv_cache_disagrees, + ); let stable = match &before { None => self.target_state(target).await.is_none(), Some(before) => self.target_state(target).await.is_some_and(|after| { - Arc::ptr_eq(before, &after) && after.active_backend_count() == count + Arc::ptr_eq(before, &after) + && after.active_backend_count() == count + && after.kv_cache_stats_disagree() == kv_cache_disagrees }), }; if stable { @@ -301,7 +320,7 @@ impl RoutingLifecycle { snapshot = rejected; let _ = self.remove_if_empty(target, target_state).await; } - self.publish_active_backend_count(target).await; + self.publish_target_metrics(target).await; } pub(super) async fn remove_inference_server_from_target( @@ -315,7 +334,7 @@ impl RoutingLifecycle { target_state.remove_backend(registration); let _ = self.remove_if_empty(target, target_state).await; - self.publish_active_backend_count(target).await; + self.publish_target_metrics(target).await; } pub(super) async fn remove_inference_server_targets( diff --git a/src/libraries/rust/stargate/crates/stargate/src/routing_state/tests.rs b/src/libraries/rust/stargate/crates/stargate/src/routing_state/tests.rs index 522f417a1..f86489110 100644 --- a/src/libraries/rust/stargate/crates/stargate/src/routing_state/tests.rs +++ b/src/libraries/rust/stargate/crates/stargate/src/routing_state/tests.rs @@ -318,6 +318,7 @@ fn shared_backend_a_stats() -> ModelStats { kv_cache_capacity_tokens: 1000, kv_cache_used_tokens: 100, kv_cache_free_tokens: 900, + kv_cache: None, num_running_queries: 11, max_engine_concurrency: 111, total_query_input_size: 1111, @@ -343,6 +344,7 @@ fn shared_backend_b_stats() -> ModelStats { kv_cache_capacity_tokens: 2000, kv_cache_used_tokens: 500, kv_cache_free_tokens: 1500, + kv_cache: None, num_running_queries: 7, max_engine_concurrency: 77, total_query_input_size: 777, @@ -1777,6 +1779,80 @@ fn cluster_backend_aggregate_dedupes_sources_and_averages_rtt() { assert_eq!(stats.stats_sources.len(), 3); } +#[test] +fn shared_cluster_reconciles_kv_snapshots_without_summing() { + let mut backend_a = backend!("kv-cluster", "backend-a", 1.0, 10.0, 5); + backend_a.stats.kv_cache = Some(stargate_proto::pb::KvCacheStats { + capacity_tokens: 1_000, + used_tokens: 250, + free_tokens: 750, + source_observed_at_unix_ms: 200, + complete: true, + }); + let cluster_state = RoutedClusterState::new(backend_a.registration.cluster_generation.clone()); + cluster_state.upsert_backend(Arc::new(backend_a.clone())); + + let mut backend_b = backend!( + backend_a.registration.cluster_generation.clone() => + "kv-cluster", "backend-b", 1.0, 10.0, 5 + ); + backend_b.stats.kv_cache = Some(stargate_proto::pb::KvCacheStats { + capacity_tokens: 2_000, + used_tokens: 500, + free_tokens: 1_500, + source_observed_at_unix_ms: 300, + complete: false, + }); + let backend_b_registration = backend_b.registration.clone(); + cluster_state.upsert_backend(Arc::new(backend_b.clone())); + + let snapshot = cluster_state + .routing_snapshot() + .expect("shared cluster should publish"); + let kv_cache = snapshot + .stats + .kv_cache + .expect("fresh complete KV snapshot should be selected"); + assert_eq!(kv_cache.capacity_tokens, 1_000); + assert_eq!(kv_cache.used_tokens, 250); + assert_eq!(snapshot.stats.kv_cache_capacity_tokens, 1_000); + assert_eq!(snapshot.stats.kv_cache_used_tokens, 250); + assert_eq!(snapshot.stats.kv_cache_free_tokens, 750); + assert!(!cluster_state.kv_cache_stats_disagree()); + + let backend_b_kv = backend_b.stats.kv_cache.as_mut().unwrap(); + backend_b_kv.complete = true; + cluster_state.upsert_backend(Arc::new(backend_b.clone())); + let snapshot = cluster_state.routing_snapshot().unwrap(); + assert_eq!(snapshot.stats.kv_cache.unwrap().capacity_tokens, 2_000); + assert!(cluster_state.kv_cache_stats_disagree()); + + let backend_b_kv = backend_b.stats.kv_cache.as_mut().unwrap(); + backend_b_kv.used_tokens = 501; + cluster_state.upsert_backend(Arc::new(backend_b.clone())); + let snapshot = cluster_state.routing_snapshot().unwrap(); + assert_eq!(snapshot.stats.kv_cache.unwrap().capacity_tokens, 1_000); + assert!(!cluster_state.kv_cache_stats_disagree()); + + let backend_b_kv = backend_b.stats.kv_cache.as_mut().unwrap(); + backend_b_kv.used_tokens = 500; + backend_b_kv.source_observed_at_unix_ms = 200; + cluster_state.upsert_backend(Arc::new(backend_b)); + let snapshot = cluster_state.routing_snapshot().unwrap(); + let kv_cache = snapshot.stats.kv_cache.unwrap(); + assert_eq!(kv_cache.capacity_tokens, 2_000); + assert_eq!(kv_cache.used_tokens, 500); + assert!(cluster_state.kv_cache_stats_disagree()); + + assert_eq!( + cluster_state.remove_backend(&backend_b_registration), + super::snapshots::ClusterBackendRemoval::Removed + ); + let snapshot = cluster_state.routing_snapshot().unwrap(); + assert_eq!(snapshot.stats.kv_cache.unwrap().capacity_tokens, 1_000); + assert!(!cluster_state.kv_cache_stats_disagree()); +} + #[tokio::test] async fn list_active_models_filters_by_routing_key() { let scenario = RegistrationScenario::new(None); From 9e1aaf98c2a8136eb3d3a8ccab5230e7a9daf54a Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Tue, 18 Aug 2026 23:18:59 +0000 Subject: [PATCH 4/9] test(stargate): mark opaque tunnel fixtures passthrough Signed-off-by: Barry Greengus --- .../stargate/crates/stargate/src/tunnel/tests.rs | 9 ++++++--- .../crates/stargate/tests/suite/proxy_contract.rs | 15 ++++++++------- 2 files changed, 14 insertions(+), 10 deletions(-) diff --git a/src/libraries/rust/stargate/crates/stargate/src/tunnel/tests.rs b/src/libraries/rust/stargate/crates/stargate/src/tunnel/tests.rs index 1ab98b137..b88a6c87d 100644 --- a/src/libraries/rust/stargate/crates/stargate/src/tunnel/tests.rs +++ b/src/libraries/rust/stargate/crates/stargate/src/tunnel/tests.rs @@ -44,7 +44,7 @@ use crate::routing_state::{ test_registration_generation, }; use pylon_lib::{ - QuicHttpTunnelConfig, ReverseQuicTunnelConfig, start_quic_http_tunnel, + QuicHttpTunnelConfig, ReverseQuicTunnelConfig, UpstreamBackend, start_quic_http_tunnel, start_reverse_quic_tunnel, }; use stargate_runtime::CriticalTaskGroup; @@ -194,6 +194,7 @@ async fn start_test_quic_tunnel( ) -> pylon_lib::QuicHttpTunnelHandle { let mut config = QuicHttpTunnelConfig::new("127.0.0.1:0".parse().unwrap(), backend_url); config.tunnel_protocol = tunnel_protocol; + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; start_quic_http_tunnel(config).await.unwrap() } @@ -539,11 +540,13 @@ fn reverse_tunnel_config( server_id: &str, backend_url: String, ) -> ReverseQuicTunnelConfig { - ReverseQuicTunnelConfig::new( + let mut config = ReverseQuicTunnelConfig::new( format!("127.0.0.1:{}", listener_addr.port()), server_id.to_string(), backend_url, - ) + ); + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; + config } struct ReverseTunnelFixture { diff --git a/src/libraries/rust/stargate/crates/stargate/tests/suite/proxy_contract.rs b/src/libraries/rust/stargate/crates/stargate/tests/suite/proxy_contract.rs index 03da3bf23..4260b6aae 100644 --- a/src/libraries/rust/stargate/crates/stargate/tests/suite/proxy_contract.rs +++ b/src/libraries/rust/stargate/crates/stargate/tests/suite/proxy_contract.rs @@ -35,7 +35,7 @@ use prometheus::{Encoder, TextEncoder}; use pylon_lib::{ CurrentModelStats, InferenceServerRegistrationClient, InferenceServerRegistrationConfig, PylonRuntimeState, QuicHttpTunnelConfig, QuicHttpTunnelHandle, RequestObservation, - RequestObservationEndpoint, RequestObservationState, TunnelTransportProtocol, + RequestObservationEndpoint, RequestObservationState, TunnelTransportProtocol, UpstreamBackend, start_quic_http_tunnel, }; use stargate::proxy::ProxyRetryConfig; @@ -671,12 +671,12 @@ async fn start_embeddings_inst( axum::serve(listener, app).await.unwrap(); }); - let tunnel = start_quic_http_tunnel(QuicHttpTunnelConfig::new( - "127.0.0.1:0".parse().unwrap(), - format!("http://{addr}"), - )) - .await - .expect("embedding tunnel failed to start"); + let mut config = + QuicHttpTunnelConfig::new("127.0.0.1:0".parse().unwrap(), format!("http://{addr}")); + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; + let tunnel = start_quic_http_tunnel(config) + .await + .expect("embedding tunnel failed to start"); let quic_url = format!("quic://{}", tunnel.listen_addr()); (addr, quic_url, tunnel, capture) } @@ -774,6 +774,7 @@ async fn start_direct_endpoint_contract_inst( let mut config = QuicHttpTunnelConfig::new("127.0.0.1:0".parse().unwrap(), format!("http://{addr}")); config.tunnel_protocol = protocol; + config.forwarding.upstream_backend = UpstreamBackend::Passthrough; let tunnel = start_quic_http_tunnel(config) .await .expect("direct endpoint contract tunnel failed to start"); From b7e37a78768c6ec49c9037904e8e935c1a575e5a Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Wed, 19 Aug 2026 20:11:30 +0000 Subject: [PATCH 5/9] refactor(pylon): derive KV snapshot completeness Signed-off-by: Barry Greengus --- .../stargate/crates/mock-dynamo/src/openai.rs | 3 +-- .../crates/pylon-lib/src/stats/collector.rs | 2 +- .../pylon-lib/src/stats/kv_stats_stream.rs | 16 +++++++--------- 3 files changed, 9 insertions(+), 12 deletions(-) diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs index 5aab886de..cb4e81287 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs @@ -418,8 +418,7 @@ pub(crate) async fn kv_cache_stats_stream(State(state): State) -> Resp "role": "aggregated", "capacity_tokens": stats.kv_cache_capacity_tokens, "used_tokens": stats.kv_cache_used_tokens, - "free_tokens": stats.kv_cache_free_tokens, - "complete": true + "free_tokens": stats.kv_cache_free_tokens }, "pools": [] }] diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs index 39da72677..650c4fad8 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs @@ -2377,7 +2377,7 @@ mod tests { #[tokio::test] async fn kv_cache_stream_updates_model_metrics() { - const SNAPSHOT: &str = "{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":42,\"models\":[{\"model\":\"model-a\",\"aliases\":[],\"routing_cache\":{\"role\":\"decode\",\"capacity_tokens\":1000,\"used_tokens\":400,\"free_tokens\":600,\"complete\":true},\"pools\":[]}]}\n"; + const SNAPSHOT: &str = "{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":42,\"models\":[{\"model\":\"model-a\",\"aliases\":[],\"routing_cache\":{\"role\":\"decode\",\"capacity_tokens\":1000,\"used_tokens\":400,\"free_tokens\":600},\"pools\":[]}]}\n"; let requests = Arc::new(AtomicUsize::new(0)); let handler_requests = requests.clone(); let metrics = PylonMetrics::new().expect("metrics should initialize"); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs index 3641cb274..d0b759111 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs @@ -42,7 +42,6 @@ struct RawRoutingCacheStats { capacity_tokens: Option, used_tokens: Option, free_tokens: Option, - complete: bool, } pub(super) fn parse_kv_stats_snapshot(line: &[u8]) -> anyhow::Result { @@ -86,7 +85,7 @@ pub(super) fn parse_kv_stats_snapshot(line: &[u8]) -> anyhow::Result 0 && used.checked_add(free) == Some(capacity) + capacity > 0 && used.checked_add(free) == Some(capacity) }); let (capacity, used, free) = routing .map(|routing| { @@ -224,9 +223,9 @@ mod tests { use super::*; #[test] - fn parses_complete_snapshot_without_placement_data() { + fn routing_cache_presence_denotes_a_complete_snapshot() { let snapshot = parse_kv_stats_snapshot( - br#"{"v":1,"type":"kv_stats_snapshot","snapshot_id":9,"observed_at_unix_ms":42,"models":[{"model":"m","aliases":["alias"],"routing_cache":{"role":"decode","capacity_tokens":100,"used_tokens":40,"free_tokens":60,"complete":true},"pools":[]}]}"#, + br#"{"v":1,"type":"kv_stats_snapshot","snapshot_id":9,"observed_at_unix_ms":42,"models":[{"model":"m","aliases":["alias"],"routing_cache":{"role":"decode","capacity_tokens":100,"used_tokens":40,"free_tokens":60},"pools":[]}]}"#, ) .unwrap(); assert_eq!(snapshot.models[0].source_observed_at_unix_ms, 42); @@ -236,9 +235,9 @@ mod tests { } #[test] - fn inconsistent_complete_snapshot_is_not_usable() { + fn inconsistent_routing_cache_snapshot_is_not_usable() { let snapshot = parse_kv_stats_snapshot( - br#"{"v":1,"type":"kv_stats_snapshot","observed_at_unix_ms":42,"models":[{"model":"m","aliases":[],"routing_cache":{"capacity_tokens":100,"used_tokens":80,"free_tokens":30,"complete":true},"pools":[]}]}"#, + br#"{"v":1,"type":"kv_stats_snapshot","observed_at_unix_ms":42,"models":[{"model":"m","aliases":[],"routing_cache":{"capacity_tokens":100,"used_tokens":80,"free_tokens":30},"pools":[]}]}"#, ) .unwrap(); assert!(!snapshot.models[0].complete); @@ -263,8 +262,7 @@ mod tests { "role": "decode", "capacity_tokens": 65_536_000, "used_tokens": 6_553_600, - "free_tokens": 58_982_400, - "complete": true + "free_tokens": 58_982_400 }, "pools": [{ "namespace": "dynamo", @@ -307,7 +305,7 @@ mod tests { b"{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":9,", )), Ok(Bytes::from_static( - b"\"models\":[{\"model\":\"m\",\"routing_cache\":{\"capacity_tokens\":10,\"used_tokens\":4,\"free_tokens\":6,\"complete\":true}}]}\n", + b"\"models\":[{\"model\":\"m\",\"routing_cache\":{\"capacity_tokens\":10,\"used_tokens\":4,\"free_tokens\":6}}]}\n", )), ]); let (tx, rx) = flume::bounded(1); From e22d1a13f6ca109caae4db4ba67d769f7251b91a Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Wed, 19 Aug 2026 21:15:04 +0000 Subject: [PATCH 6/9] feat(pylon): consume unified Dynamo stats over gRPC Signed-off-by: Barry Greengus --- src/libraries/rust/stargate/Cargo.lock | 209 +-- .../stargate/crates/mock-dynamo/BUILD.bazel | 4 +- .../stargate/crates/mock-dynamo/Cargo.toml | 3 + .../stargate/crates/mock-dynamo/src/main.rs | 18 +- .../stargate/crates/mock-dynamo/src/openai.rs | 55 +- .../crates/mock-dynamo/src/stats_stream.rs | 215 ++- .../crates/mock-dynamo/src/test_control.rs | 8 +- .../stargate/crates/mock-dynamo/src/tests.rs | 58 +- .../proto/proto/dynamo_frontend_stats.proto | 219 +++ .../stargate/crates/proto/src/build_plan.rs | 22 +- .../rust/stargate/crates/proto/src/lib.rs | 13 +- .../rust/stargate/crates/pylon-lib/Cargo.toml | 5 - .../pylon-lib/benches/engine_stats_stream.rs | 318 ---- .../rust/stargate/crates/pylon-lib/src/lib.rs | 5 +- .../crates/pylon-lib/src/model_lifecycle.rs | 106 +- .../crates/pylon-lib/src/stats/aggregator.rs | 22 +- .../crates/pylon-lib/src/stats/collector.rs | 134 +- .../src/stats/engine_stats_stream.rs | 1388 ++++------------- .../crates/pylon-lib/src/stats/kv_stats.rs | 113 ++ .../pylon-lib/src/stats/kv_stats_stream.rs | 377 ----- .../crates/pylon-lib/src/stats/mod.rs | 4 +- .../rust/stargate/crates/pylon/src/main.rs | 40 +- .../rust/stargate/crates/pylon/src/startup.rs | 38 +- .../crates/stargate-bench/src/k8s/render.rs | 2 +- .../crates/stargate-bench/src/orchestrator.rs | 1 - .../stargate/tests/suite/integration.rs | 141 +- .../stargate/docs/diagrams/system-dfd.puml | 2 +- 27 files changed, 997 insertions(+), 2523 deletions(-) create mode 100644 src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto delete mode 100644 src/libraries/rust/stargate/crates/pylon-lib/benches/engine_stats_stream.rs create mode 100644 src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs delete mode 100644 src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs diff --git a/src/libraries/rust/stargate/Cargo.lock b/src/libraries/rust/stargate/Cargo.lock index 86a20a9ca..8c6e9dfaa 100644 --- a/src/libraries/rust/stargate/Cargo.lock +++ b/src/libraries/rust/stargate/Cargo.lock @@ -30,12 +30,6 @@ version = "0.2.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "683d7910e743518b0e34f1186f92494becacb047c7b6bf616c96772180fef923" -[[package]] -name = "anes" -version = "0.1.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299" - [[package]] name = "anstream" version = "1.0.0" @@ -371,12 +365,6 @@ dependencies = [ "capnp", ] -[[package]] -name = "cast" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5" - [[package]] name = "cc" version = "1.2.56" @@ -407,33 +395,6 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" -[[package]] -name = "ciborium" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e" -dependencies = [ - "ciborium-io", - "ciborium-ll", - "serde", -] - -[[package]] -name = "ciborium-io" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757" - -[[package]] -name = "ciborium-ll" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9" -dependencies = [ - "ciborium-io", - "half", -] - [[package]] name = "clap" version = "4.6.1" @@ -533,73 +494,12 @@ dependencies = [ "libc", ] -[[package]] -name = "criterion" -version = "0.5.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2b12d017a929603d80db1831cd3a24082f8137ce19c69e6447f54f5fc8d692f" -dependencies = [ - "anes", - "cast", - "ciborium", - "clap", - "criterion-plot", - "is-terminal", - "itertools 0.10.5", - "num-traits", - "once_cell", - "oorandom", - "plotters", - "rayon", - "regex", - "serde", - "serde_derive", - "serde_json", - "tinytemplate", - "walkdir", -] - -[[package]] -name = "criterion-plot" -version = "0.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6b50826342786a51a89e2da3a28f1c32b06e387201bc2d19791f622c673706b1" -dependencies = [ - "cast", - "itertools 0.10.5", -] - -[[package]] -name = "crossbeam-deque" -version = "0.8.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9dd111b7b7f7d55b72c0a6ae361660ee5853c9af73f70c3c2ef6858b950e2e51" -dependencies = [ - "crossbeam-epoch", - "crossbeam-utils", -] - -[[package]] -name = "crossbeam-epoch" -version = "0.9.18" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" -dependencies = [ - "crossbeam-utils", -] - [[package]] name = "crossbeam-utils" version = "0.8.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" -[[package]] -name = "crunchy" -version = "0.2.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" - [[package]] name = "crypto-common" version = "0.1.7" @@ -1092,17 +992,6 @@ dependencies = [ "tokio-util", ] -[[package]] -name = "half" -version = "2.7.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" -dependencies = [ - "cfg-if", - "crunchy", - "zerocopy", -] - [[package]] name = "hashbrown" version = "0.15.5" @@ -1135,12 +1024,6 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" -[[package]] -name = "hermit-abi" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" - [[package]] name = "hickory-proto" version = "0.24.4" @@ -1466,32 +1349,12 @@ dependencies = [ "serde", ] -[[package]] -name = "is-terminal" -version = "0.4.17" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" -dependencies = [ - "hermit-abi", - "libc", - "windows-sys 0.61.2", -] - [[package]] name = "is_terminal_polyfill" version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" -[[package]] -name = "itertools" -version = "0.10.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b0fd2260e829bddf4cb6ea802289de2f86d6a7a690192fbe91b3f46e0f2c8473" -dependencies = [ - "either", -] - [[package]] name = "itertools" version = "0.14.0" @@ -1860,9 +1723,12 @@ dependencies = [ "async-stream", "axum", "clap", + "futures", "serde", "serde_json", + "stargate-proto", "tokio", + "tonic", "tracing", "tracing-subscriber", ] @@ -1967,12 +1833,6 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" -[[package]] -name = "oorandom" -version = "11.1.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e" - [[package]] name = "openssl-probe" version = "0.2.1" @@ -2193,34 +2053,6 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" -[[package]] -name = "plotters" -version = "0.3.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747" -dependencies = [ - "num-traits", - "plotters-backend", - "plotters-svg", - "wasm-bindgen", - "web-sys", -] - -[[package]] -name = "plotters-backend" -version = "0.3.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a" - -[[package]] -name = "plotters-svg" -version = "0.3.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670" -dependencies = [ - "plotters-backend", -] - [[package]] name = "portable-atomic" version = "1.13.1" @@ -2330,7 +2162,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "343d3bd7056eda839b03204e68deff7d1b13aba7af2b2fd16890697274262ee7" dependencies = [ "heck", - "itertools 0.14.0", + "itertools", "log", "multimap", "petgraph", @@ -2351,7 +2183,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b570b25f7617e43d59005d0990ccb79e950a423952cea19671b7a876da390adf" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools", "proc-macro2", "quote", "syn", @@ -2461,7 +2293,6 @@ dependencies = [ "async-stream", "axum", "bytes", - "criterion", "flume", "futures", "h3", @@ -2657,26 +2488,6 @@ dependencies = [ "rand_core 0.9.5", ] -[[package]] -name = "rayon" -version = "1.12.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fb39b166781f92d482534ef4b4b1b2568f42613b53e5b6c160e24cfbfa30926d" -dependencies = [ - "either", - "rayon-core", -] - -[[package]] -name = "rayon-core" -version = "1.13.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91" -dependencies = [ - "crossbeam-deque", - "crossbeam-utils", -] - [[package]] name = "rcgen" version = "0.13.2" @@ -3681,16 +3492,6 @@ dependencies = [ "zerovec", ] -[[package]] -name = "tinytemplate" -version = "1.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc" -dependencies = [ - "serde", - "serde_json", -] - [[package]] name = "tinyvec" version = "1.10.0" diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/BUILD.bazel b/src/libraries/rust/stargate/crates/mock-dynamo/BUILD.bazel index 05a653501..60d47e298 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/BUILD.bazel +++ b/src/libraries/rust/stargate/crates/mock-dynamo/BUILD.bazel @@ -15,5 +15,7 @@ rust_binary( edition = "2024", proc_macro_deps = all_crate_deps(proc_macro = True), visibility = ["//visibility:public"], - deps = all_crate_deps(normal = True), + deps = [ + "//src/libraries/rust/stargate/crates/proto:stargate-proto", + ] + all_crate_deps(normal = True), ) diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/Cargo.toml b/src/libraries/rust/stargate/crates/mock-dynamo/Cargo.toml index 5534b63b5..0e41cb026 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/Cargo.toml +++ b/src/libraries/rust/stargate/crates/mock-dynamo/Cargo.toml @@ -35,8 +35,11 @@ anyhow = { workspace = true } async-stream = { workspace = true } axum = { workspace = true } clap = { workspace = true } +futures = { workspace = true } serde = { workspace = true } serde_json = { workspace = true } +stargate-proto = { workspace = true } tokio = { workspace = true } +tonic = { workspace = true } tracing = { workspace = true } tracing-subscriber = { workspace = true } diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs index 654e4a99a..c344d6c84 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs @@ -85,7 +85,7 @@ struct AppState { health_delay: Duration, kv_cache: Arc>, stats_events: broadcast::Sender, - kv_stats_enabled: Arc, + stats_stream_enabled: Arc, test_control: test_control::TestControlState, } @@ -118,20 +118,17 @@ async fn main() -> Result<()> { args.kv_cache_capacity_tokens, ))), stats_events, - kv_stats_enabled: Arc::new(AtomicBool::new(true)), + stats_stream_enabled: Arc::new(AtomicBool::new(true)), test_control: test_control::TestControlState::with_discovered_models([args.model_name]), }; + let grpc = stats_stream::grpc_router(state.clone()); let app = Router::new() .route("/v1/chat/completions", post(openai::chat_completions)) .route("/v1/models", get(openai::list_models)) .route("/v1/responses", post(openai::responses)) .route("/v1/embeddings", post(openai::embeddings)) - .route("/pylon/v1/stats/stream", get(stats_stream::stats_stream)) - .route( - "/v1/kv-cache/stats/stream", - get(openai::kv_cache_stats_stream), - ) + .route("/v1/stats/stream", get(stats_stream::stats_stream)) .route("/kv-cache/stats", get(openai::kv_cache_stats)) .route( "/test-control/models/{model}", @@ -139,15 +136,16 @@ async fn main() -> Result<()> { ) .route("/test-control", get(test_control::test_control_snapshot)) .route( - "/test-control/kv-stats", - put(test_control::update_kv_stats_test_control), + "/test-control/stats-stream", + put(test_control::update_stats_stream_test_control), ) .route( "/test-control/discovery-models", put(test_control::replace_discovery_models), ) .route("/health", get(openai::health)) - .with_state(state); + .with_state(state) + .merge(grpc); let listener = TcpListener::bind(http_addr).await?; let actual_http_addr = listener.local_addr()?; diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs index cb4e81287..61d637cf5 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs @@ -14,14 +14,11 @@ // limitations under the License. use axum::Json; -use axum::body::{Body, Bytes}; use axum::extract::State; -use axum::http::{HeaderMap, HeaderValue, StatusCode, header}; +use axum::http::{HeaderMap, StatusCode}; use axum::response::sse::{Event, KeepAlive, Sse}; use axum::response::{IntoResponse, Response}; use serde::{Deserialize, Serialize}; -use std::convert::Infallible; -use std::sync::atomic::Ordering; use std::time::{Duration, SystemTime, UNIX_EPOCH}; use tokio::sync::OwnedSemaphorePermit; use tracing::info; @@ -392,51 +389,6 @@ pub(crate) async fn kv_cache_stats(State(state): State) -> Json) -> Response { - if !state.kv_stats_enabled.load(Ordering::Relaxed) { - return StatusCode::SERVICE_UNAVAILABLE.into_response(); - } - let stream = async_stream::stream! { - let mut interval = tokio::time::interval(Duration::from_secs(1)); - interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); - let mut snapshot_id = 1_u64; - loop { - interval.tick().await; - if !state.kv_stats_enabled.load(Ordering::Relaxed) { - break; - } - let stats = state.kv_cache.lock().await.stats(&state.model_name); - let snapshot = serde_json::json!({ - "v": 1, - "type": "kv_stats_snapshot", - "snapshot_id": snapshot_id, - "observed_at_unix_ms": unix_millis(), - "models": [{ - "model": stats.model, - "aliases": [], - "routing_cache": { - "role": "aggregated", - "capacity_tokens": stats.kv_cache_capacity_tokens, - "used_tokens": stats.kv_cache_used_tokens, - "free_tokens": stats.kv_cache_free_tokens - }, - "pools": [] - }] - }); - snapshot_id = snapshot_id.saturating_add(1); - let mut line = serde_json::to_vec(&snapshot).expect("mock KV stats serialize"); - line.push(b'\n'); - yield Ok::(Bytes::from(line)); - } - }; - let mut response = Response::new(Body::from_stream(stream)); - response.headers_mut().insert( - header::CONTENT_TYPE, - HeaderValue::from_static("application/x-ndjson"), - ); - response -} - impl AppState { async fn process_input_with_cache( &self, @@ -483,8 +435,7 @@ impl AppState { output_tokens: impl Into>, finished: bool, ) { - let _ = self.stats_events.send(StatsStreamEvent::Stats { - v: 1, + let _ = self.stats_events.send(StatsStreamEvent { request_id: request_id.to_string(), model: model.to_string(), tokens_processed: Some(input_tokens as u64), @@ -667,6 +618,6 @@ fn current_unix_timestamp() -> u64 { time_since_epoch().as_secs() } -fn unix_millis() -> u64 { +pub(crate) fn unix_millis() -> u64 { u64::try_from(time_since_epoch().as_millis()).unwrap_or(u64::MAX) } diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs index 5d869c4e6..73b5cfa40 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs @@ -1,77 +1,184 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. + +use std::convert::Infallible; +use std::pin::Pin; +use std::sync::atomic::Ordering; +use std::time::Duration; use axum::body::{Body, Bytes}; use axum::extract::State; -use axum::response::Response; -use serde::Serialize; -use std::time::Duration; +use axum::http::{HeaderValue, StatusCode, header}; +use axum::response::{IntoResponse, Response}; +use futures::Stream; +use stargate_proto::dynamo_frontend_stats as proto; use tokio::sync::broadcast; +use tonic::{Request, Response as GrpcResponse, Status}; use crate::AppState; -#[derive(Debug, Clone, Serialize)] -#[serde(tag = "type")] -pub(crate) enum StatsStreamEvent { - #[serde(rename = "stats")] - Stats { - v: u8, - request_id: String, - model: String, - #[serde(skip_serializing_if = "Option::is_none")] - tokens_processed: Option, - #[serde(skip_serializing_if = "Option::is_none")] - tokens_generated: Option, - #[serde(skip_serializing_if = "is_false")] - finished: bool, - }, - #[serde(rename = "ping")] - Ping { v: u8 }, +#[derive(Debug, Clone)] +pub(crate) struct StatsStreamEvent { + pub(crate) request_id: String, + pub(crate) model: String, + pub(crate) tokens_processed: Option, + pub(crate) tokens_generated: Option, + pub(crate) finished: bool, } -fn is_false(value: &bool) -> bool { - !*value +pub(crate) fn grpc_router(state: AppState) -> axum::Router { + let service = + proto::frontend_stats_server::FrontendStatsServer::new(MockFrontendStats { state }); + tonic::service::Routes::new(service).into_axum_router() } pub(crate) async fn stats_stream(State(state): State) -> Response { - let mut events = state.stats_events.subscribe(); + if !state.stats_stream_enabled.load(Ordering::Relaxed) { + return StatusCode::SERVICE_UNAVAILABLE.into_response(); + } + let source = mixed_stats_stream(state); let stream = async_stream::stream! { - let mut ping = tokio::time::interval(Duration::from_secs(1)); - ping.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + futures::pin_mut!(source); + while let Some(update) = futures::StreamExt::next(&mut source).await { + yield Ok::(ndjson_event(update)); + } + }; + let mut response = Response::new(Body::from_stream(stream)); + response.headers_mut().insert( + header::CONTENT_TYPE, + HeaderValue::from_static("application/x-ndjson"), + ); + response +} + +fn mixed_stats_stream(state: AppState) -> impl Stream + Send + 'static { + let mut events = state.stats_events.subscribe(); + async_stream::stream! { + let mut snapshots = tokio::time::interval(Duration::from_secs(1)); + snapshots.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + let mut snapshot_id = 1_u64; loop { - let event = tokio::select! { - event = events.recv() => { - match event { - Ok(event) => event, - Err(broadcast::error::RecvError::Lagged(_)) => continue, - Err(broadcast::error::RecvError::Closed) => break, - } + if !state.stats_stream_enabled.load(Ordering::Relaxed) { + break; + } + tokio::select! { + event = events.recv() => match event { + Ok(event) => yield request_update(event), + Err(broadcast::error::RecvError::Lagged(_)) => break, + Err(broadcast::error::RecvError::Closed) => break, + }, + _ = snapshots.tick() => { + let stats = state.kv_cache.lock().await.stats(&state.model_name); + yield kv_update(snapshot_id, stats); + snapshot_id = snapshot_id.saturating_add(1); } - _ = ping.tick() => StatsStreamEvent::Ping { v: 1 }, - }; - yield Ok::(ndjson_event(&event)); + } } - }; - Response::builder() - .header("content-type", "application/x-ndjson") - .body(Body::from_stream(stream)) - .expect("stats stream response should build") + } +} + +fn request_update(event: StatsStreamEvent) -> proto::StatsUpdate { + proto::StatsUpdate { + update: Some(proto::stats_update::Update::RequestStats( + proto::RequestStats { + request_id: event.request_id, + model: event.model, + tokens_processed: event.tokens_processed, + tokens_generated: event.tokens_generated, + finished: event.finished, + }, + )), + } +} + +fn kv_update(snapshot_id: u64, stats: crate::kv_cache::KvCacheStats) -> proto::StatsUpdate { + proto::StatsUpdate { + update: Some(proto::stats_update::Update::KvStats( + proto::KvStatsSnapshot { + snapshot_id, + observed_at_unix_ms: crate::openai::unix_millis(), + models: vec![proto::ModelKvStats { + model: stats.model, + aliases: Vec::new(), + routing_cache: Some(proto::RoutingCacheStats { + role: proto::WorkerRole::Aggregated as i32, + capacity_tokens: stats.kv_cache_capacity_tokens, + used_tokens: stats.kv_cache_used_tokens, + free_tokens: stats.kv_cache_free_tokens, + }), + pools: Vec::new(), + }], + }, + )), + } } -pub(crate) fn ndjson_event(event: &StatsStreamEvent) -> Bytes { - let mut line = serde_json::to_vec(event).expect("stats stream event should serialize"); +pub(crate) fn ndjson_event(update: proto::StatsUpdate) -> Bytes { + let value = match update + .update + .expect("mock stats update must have a payload") + { + proto::stats_update::Update::RequestStats(event) => serde_json::json!({ + "v": 1, + "type": "stats", + "request_id": event.request_id, + "model": event.model, + "tokens_processed": event.tokens_processed, + "tokens_generated": event.tokens_generated, + "finished": event.finished, + }), + proto::stats_update::Update::KvStats(snapshot) => serde_json::json!({ + "v": 1, + "type": "kv_stats_snapshot", + "snapshot_id": snapshot.snapshot_id, + "observed_at_unix_ms": snapshot.observed_at_unix_ms, + "models": snapshot.models.into_iter().map(|model| serde_json::json!({ + "model": model.model, + "aliases": model.aliases, + "routing_cache": model.routing_cache.map(|cache| serde_json::json!({ + "role": "aggregated", + "capacity_tokens": cache.capacity_tokens, + "used_tokens": cache.used_tokens, + "free_tokens": cache.free_tokens, + })), + "pools": [], + })).collect::>(), + }), + }; + let mut line = serde_json::to_vec(&value).expect("mock stats update should serialize"); line.push(b'\n'); Bytes::from(line) } + +#[derive(Clone)] +struct MockFrontendStats { + state: AppState, +} + +#[tonic::async_trait] +impl proto::frontend_stats_server::FrontendStats for MockFrontendStats { + type WatchStatsStream = + Pin> + Send + 'static>>; + type WatchKvPlacementsStream = + Pin> + Send + 'static>>; + + async fn watch_stats( + &self, + _request: Request, + ) -> Result, Status> { + if !self.state.stats_stream_enabled.load(Ordering::Relaxed) { + return Err(Status::unavailable("mock stats stream disabled")); + } + let stream = futures::StreamExt::map(mixed_stats_stream(self.state.clone()), Ok); + Ok(GrpcResponse::new(Box::pin(stream))) + } + + async fn watch_kv_placements( + &self, + _request: Request, + ) -> Result, Status> { + Err(Status::unimplemented( + "mock Dynamo does not model KV placements", + )) + } +} diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs index 31f6912ea..79e7239a9 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/test_control.rs @@ -254,16 +254,16 @@ pub(crate) async fn test_control_snapshot( } #[derive(Debug, Clone, Deserialize)] -pub(crate) struct KvStatsTestControlUpdate { +pub(crate) struct StatsStreamTestControlUpdate { pub(crate) enabled: bool, } -pub(crate) async fn update_kv_stats_test_control( +pub(crate) async fn update_stats_stream_test_control( State(state): State, - Json(update): Json, + Json(update): Json, ) -> axum::http::StatusCode { state - .kv_stats_enabled + .stats_stream_enabled .store(update.enabled, Ordering::Relaxed); axum::http::StatusCode::NO_CONTENT } diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs index 0126d54ce..032582602 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs @@ -53,23 +53,23 @@ fn test_state() -> AppState { health_delay: Duration::ZERO, kv_cache: Arc::new(Mutex::new(KvCacheState::new(0))), stats_events: test_stats_events(), - kv_stats_enabled: Arc::new(AtomicBool::new(true)), + stats_stream_enabled: Arc::new(AtomicBool::new(true)), test_control: TestControlState::with_discovered_models(["dummy-model".to_string()]), } } #[tokio::test] -async fn kv_stats_test_control_does_not_disable_health() { +async fn stats_stream_test_control_does_not_disable_health() { let state = test_state(); - let status = update_kv_stats_test_control( + let status = update_stats_stream_test_control( State(state.clone()), - Json(KvStatsTestControlUpdate { enabled: false }), + Json(StatsStreamTestControlUpdate { enabled: false }), ) .await; assert_eq!(status, axum::http::StatusCode::NO_CONTENT); assert_eq!( - kv_cache_stats_stream(State(state.clone())).await.status(), + stats_stream(State(state.clone())).await.status(), axum::http::StatusCode::SERVICE_UNAVAILABLE ); assert_eq!(health(State(state)).await, "ok"); @@ -675,7 +675,7 @@ async fn streaming_response_delays_first_data_frame_until_ttft() { }; let app = Router::new() .route("/v1/chat/completions", post(chat_completions)) - .route("/pylon/v1/stats/stream", get(stats_stream)) + .route("/v1/stats/stream", get(stats_stream)) .with_state(state); let (addr, server) = spawn_test_app(app).await; @@ -710,7 +710,7 @@ async fn streaming_response_exposes_stats_stream_endpoint() { }; let app = Router::new() .route("/v1/chat/completions", post(chat_completions)) - .route("/pylon/v1/stats/stream", get(stats_stream)) + .route("/v1/stats/stream", get(stats_stream)) .with_state(state); let (addr, server) = spawn_test_app(app).await; @@ -732,10 +732,8 @@ async fn streaming_response_exposes_stats_stream_endpoint() { .expect("test client should connect"); stream .write_all( - format!( - "GET /pylon/v1/stats/stream HTTP/1.1\r\nhost: {addr}\r\nconnection: close\r\n\r\n", - ) - .as_bytes(), + format!("GET /v1/stats/stream HTTP/1.1\r\nhost: {addr}\r\nconnection: close\r\n\r\n",) + .as_bytes(), ) .await .expect("stats stream request should write"); @@ -754,23 +752,35 @@ async fn streaming_response_exposes_stats_stream_endpoint() { #[test] fn stats_stream_events_are_ndjson() { - let event = StatsStreamEvent::Stats { - v: 1, - request_id: "req-1".to_string(), - model: "dummy-model".to_string(), - tokens_processed: Some(11), - tokens_generated: Some(2), - finished: true, + let event = stargate_proto::dynamo_frontend_stats::StatsUpdate { + update: Some( + stargate_proto::dynamo_frontend_stats::stats_update::Update::RequestStats( + stargate_proto::dynamo_frontend_stats::RequestStats { + request_id: "req-1".to_string(), + model: "dummy-model".to_string(), + tokens_processed: Some(11), + tokens_generated: Some(2), + finished: true, + }, + ), + ), }; - let line = String::from_utf8(ndjson_event(&event).to_vec()).unwrap(); + let line = ndjson_event(event); + let value: serde_json::Value = serde_json::from_slice(&line).unwrap(); assert_eq!( - line, - r#"{"type":"stats","v":1,"request_id":"req-1","model":"dummy-model","tokens_processed":11,"tokens_generated":2,"finished":true}"# - .to_string() - + "\n" - ); + value, + serde_json::json!({ + "v": 1, + "type": "stats", + "request_id": "req-1", + "model": "dummy-model", + "tokens_processed": 11, + "tokens_generated": 2, + "finished": true, + }) + ); } #[tokio::test] diff --git a/src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto b/src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto new file mode 100644 index 000000000..cf5165a59 --- /dev/null +++ b/src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto @@ -0,0 +1,219 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. +// All rights reserved. SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +syntax = "proto3"; + +package dynamo.frontend.stats.v1; + +service FrontendStats +{ + rpc WatchStats(WatchStatsRequest) returns (stream StatsUpdate); + rpc WatchKvPlacements(WatchKvPlacementsRequest) returns (stream KvPlacementUpdate); +} + +message WatchStatsRequest {} + +message StatsUpdate +{ + oneof update + { + RequestStats request_stats = 1; + KvStatsSnapshot kv_stats = 2; + } +} + +message RequestStats +{ + string request_id = 1; + string model = 2; + optional uint64 tokens_processed = 3; + optional uint64 tokens_generated = 4; + bool finished = 5; +} + +message KvStatsSnapshot +{ + uint64 snapshot_id = 1; + uint64 observed_at_unix_ms = 2; + repeated ModelKvStats models = 3; +} + +message ModelKvStats +{ + string model = 1; + repeated string aliases = 2; + optional RoutingCacheStats routing_cache = 3; + repeated KvPoolStats pools = 4; +} + +enum WorkerRole { + WORKER_ROLE_UNSPECIFIED = 0; + WORKER_ROLE_AGGREGATED = 1; + WORKER_ROLE_PREFILL = 2; + WORKER_ROLE_DECODE = 3; + WORKER_ROLE_ENCODE = 4; +} + +message RoutingCacheStats +{ + WorkerRole role = 1; + uint64 capacity_tokens = 2; + uint64 used_tokens = 3; + uint64 free_tokens = 4; +} + +message KvPoolStats +{ + string namespace = 1; + string component = 2; + string endpoint = 3; + WorkerRole role = 4; + StorageTier storage_tier = 5; + uint32 block_size_tokens = 6; + uint64 expected_ranks = 7; + uint64 observed_ranks = 8; + optional uint64 capacity_blocks = 9; + optional uint64 used_blocks = 10; + optional uint64 free_blocks = 11; + optional uint64 active_decode_blocks = 12; + optional uint64 active_prefill_tokens = 13; + bool complete = 14; +} + +message WatchKvPlacementsRequest {} + +message KvPlacementUpdate +{ + oneof update + { + KvPlacementSnapshotBoundary snapshot_begin = 1; + KvPlacementEvents snapshot_events = 2; + KvPlacementSnapshotBoundary snapshot_end = 3; + KvPlacementEvents events = 4; + KvPlacementSourceError source_error = 5; + } +} + +message KvPlacementSnapshotBoundary +{ + uint64 snapshot_id = 1; + bool complete = 2; + repeated KvPlacementCursor cursors = 3; +} + +message KvPlacementCursor +{ + string model = 1; + string namespace = 2; + string component = 3; + string endpoint = 4; + uint64 cursor = 5; +} + +message KvPlacementSource +{ + string model = 1; + string namespace = 2; + string component = 3; + string endpoint = 4; + uint32 block_size_tokens = 5; +} + +message KvPlacementEvents +{ + optional uint64 snapshot_id = 1; + KvPlacementSource source = 2; + uint64 cursor = 3; + uint32 batch_index = 4; + uint32 batch_count = 5; + repeated RouterEvent events = 6; +} + +message KvPlacementSourceError +{ + KvPlacementSource source = 1; + string reason = 2; +} + +enum StorageTier { + STORAGE_TIER_UNSPECIFIED = 0; + STORAGE_TIER_DEVICE = 1; + STORAGE_TIER_HOST_PINNED = 2; + STORAGE_TIER_DISK = 3; + STORAGE_TIER_EXTERNAL = 4; +} + +message ResidencyDomain +{ + enum Kind { + KIND_MISSING = 0; + KIND_WORKER = 1; + KIND_CACHE_OWNER = 2; + KIND_UNKNOWN = 3; + KIND_INVALID = 4; + } + + Kind kind = 1; + string unknown_value = 2; +} + +message RouterEvent +{ + uint64 worker_id = 1; + StorageTier storage_tier = 2; + ResidencyDomain residency_domain = 3; + optional string state_source = 4; + uint64 event_id = 5; + uint32 dp_rank = 6; + oneof data + { + KvCacheStore stored = 7; + KvCacheRemove removed = 8; + KvCacheClear cleared = 9; + } +} + +message KvCacheStore +{ + optional uint64 parent_hash = 1; + optional uint32 start_position = 2; + repeated KvCacheBlock blocks = 3; +} + +message KvCacheBlock +{ + uint64 block_hash = 1; + uint64 tokens_hash = 2; + repeated MultimodalObject multimodal_objects = 3; +} + +message MultimodalObject +{ + uint64 hash = 1; + repeated TokenRange offsets = 2; +} + +message TokenRange +{ + uint64 start = 1; + uint64 end = 2; +} + +message KvCacheRemove +{ + repeated uint64 block_hashes = 1; +} + +message KvCacheClear {} diff --git a/src/libraries/rust/stargate/crates/proto/src/build_plan.rs b/src/libraries/rust/stargate/crates/proto/src/build_plan.rs index 3c26fd49f..1863c81f6 100644 --- a/src/libraries/rust/stargate/crates/proto/src/build_plan.rs +++ b/src/libraries/rust/stargate/crates/proto/src/build_plan.rs @@ -24,7 +24,7 @@ pub(crate) struct ProtoCompilePlan { pub field_attributes: &'static [(&'static str, &'static str)], } -pub(crate) fn proto_compile_plans() -> [ProtoCompilePlan; 2] { +pub(crate) fn proto_compile_plans() -> [ProtoCompilePlan; 3] { [ ProtoCompilePlan { protos: &["proto/stargate.proto"], @@ -49,6 +49,13 @@ pub(crate) fn proto_compile_plans() -> [ProtoCompilePlan; 2] { type_attributes: &[], field_attributes: &[], }, + ProtoCompilePlan { + protos: &["proto/dynamo_frontend_stats.proto"], + includes: &["proto"], + build_server: true, + type_attributes: &[], + field_attributes: &[], + }, ] } @@ -58,7 +65,7 @@ mod tests { #[test] fn stargate_plan_carries_watch_response_serde_attributes() { - let [stargate_plan, _] = proto_compile_plans(); + let [stargate_plan, _, _] = proto_compile_plans(); assert_eq!(stargate_plan.protos, ["proto/stargate.proto"]); assert!(stargate_plan.build_server); @@ -68,7 +75,7 @@ mod tests { #[test] fn gateway_plan_builds_client_and_server_proto() { - let [_, gateway_plan] = proto_compile_plans(); + let [_, gateway_plan, _] = proto_compile_plans(); assert_eq!(gateway_plan.protos, ["proto/llm_gateway.proto"]); assert_eq!(gateway_plan.includes, ["proto"]); @@ -76,4 +83,13 @@ mod tests { assert!(gateway_plan.type_attributes.is_empty()); assert!(gateway_plan.field_attributes.is_empty()); } + + #[test] + fn dynamo_frontend_stats_plan_builds_client_and_server_proto() { + let [_, _, stats_plan] = proto_compile_plans(); + + assert_eq!(stats_plan.protos, ["proto/dynamo_frontend_stats.proto"]); + assert_eq!(stats_plan.includes, ["proto"]); + assert!(stats_plan.build_server); + } } diff --git a/src/libraries/rust/stargate/crates/proto/src/lib.rs b/src/libraries/rust/stargate/crates/proto/src/lib.rs index 5be9ff315..7ba6cf200 100644 --- a/src/libraries/rust/stargate/crates/proto/src/lib.rs +++ b/src/libraries/rust/stargate/crates/proto/src/lib.rs @@ -17,6 +17,10 @@ pub mod gateway_pb { tonic::include_proto!("llm_gateway"); } +pub mod dynamo_frontend_stats { + tonic::include_proto!("dynamo.frontend.stats.v1"); +} + pub const REGISTRATION_HEARTBEAT_MS_METADATA: &str = "x-stargate-registration-heartbeat-ms"; pub mod pb { @@ -102,10 +106,10 @@ mod tests { use crate::pb::{StargateInfo, WatchStargatesResponse}; #[test] - fn proto_build_plan_covers_stargate_and_gateway_generation() { + fn proto_build_plan_covers_all_generation() { let plans = proto_compile_plans(); - assert_eq!(plans.len(), 2); + assert_eq!(plans.len(), 3); assert_eq!(plans[0].protos, ["proto/stargate.proto"]); assert_eq!(plans[0].includes, ["proto"]); assert!(plans[0].build_server); @@ -128,7 +132,10 @@ mod tests { assert!(plans[1].build_server); assert!(plans[1].type_attributes.is_empty()); assert!(plans[1].field_attributes.is_empty()); - assert_eq!(crate::build_script::planned_proto_compile_count(), 2); + assert_eq!(plans[2].protos, ["proto/dynamo_frontend_stats.proto"]); + assert_eq!(plans[2].includes, ["proto"]); + assert!(plans[2].build_server); + assert_eq!(crate::build_script::planned_proto_compile_count(), 3); } #[test] diff --git a/src/libraries/rust/stargate/crates/pylon-lib/Cargo.toml b/src/libraries/rust/stargate/crates/pylon-lib/Cargo.toml index 8716cca23..2e07d59c2 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/Cargo.toml +++ b/src/libraries/rust/stargate/crates/pylon-lib/Cargo.toml @@ -65,13 +65,8 @@ stargate-tls = { workspace = true } uuid = { workspace = true } [dev-dependencies] -criterion = "0.5" opentelemetry_sdk = { workspace = true } rcgen = { workspace = true } tempfile = { workspace = true } tokio = { workspace = true, features = ["test-util"] } tower = { workspace = true } - -[[bench]] -name = "engine_stats_stream" -harness = false diff --git a/src/libraries/rust/stargate/crates/pylon-lib/benches/engine_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/benches/engine_stats_stream.rs deleted file mode 100644 index f7ba4a3c1..000000000 --- a/src/libraries/rust/stargate/crates/pylon-lib/benches/engine_stats_stream.rs +++ /dev/null @@ -1,318 +0,0 @@ -// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -// SPDX-License-Identifier: Apache-2.0 -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -use std::time::{Duration, Instant}; -use std::{convert::Infallible, sync::Arc}; - -use axum::{ - Router, - body::{Body, Bytes}, - extract::State, - response::Response, - routing::get, -}; -use criterion::{Criterion, Throughput, black_box, criterion_group, criterion_main}; -use futures::{StreamExt, stream}; -use pylon_lib::{ - EngineStatsStreamConfig, EngineStatsStreamMode, PylonRuntimeState, RequestCounterUpdate, - RequestCounterUpdateInput, StatsAggregatorUpdate, StatsCollectorConfig, StatsCollectorHandle, - StatsUpdateSource, parse_engine_stats_line_for_benchmark, start_engine_stats_stream, - start_stats_collector_with_engine_stats, stats_aggregator_update_channel, -}; -use tokio::net::TcpListener; -use tokio::task::JoinHandle; -use tokio::time::Instant as TokioInstant; - -const EVENT_COUNT: u64 = 50_000; -const TEST_EVENT_COUNT: u64 = 1_024; -const TEST_SENTINEL_TIMEOUT: Duration = Duration::from_secs(10); -const REQUEST_IDS: u64 = 1_024; -const SENTINEL_OUTPUT_TOKENS: u64 = 10_000; -const SENTINEL_OUTPUT_TPS: f64 = SENTINEL_OUTPUT_TOKENS as f64; -const COMPACT_STATS_EVENT: &[u8] = br#"{"v":1,"type":"stats","request_id":"req-1","model":"model-a","tokens_processed":4096,"tokens_generated":128}"#; - -fn bench_engine_stats_stream(c: &mut Criterion) { - let test_mode = running_in_criterion_test_mode(); - let (event_count, scale) = if test_mode { - (TEST_EVENT_COUNT, "smoke") - } else { - (EVENT_COUNT, "50k") - }; - let collector_benchmark_name = - format!("collector_channel_{scale}_request_counters_to_final_snapshot"); - let endpoint_benchmark_name = - format!("http_endpoint_to_collector_{scale}_request_counters_to_final_snapshot"); - - let mut parser = c.benchmark_group("engine_stats_stream_parser"); - parser.throughput(Throughput::Elements(1)); - parser.bench_function("parse_compact_request_counter", |b| { - b.iter(|| { - parse_engine_stats_line_for_benchmark( - black_box(COMPACT_STATS_EVENT), - black_box(TokioInstant::now()), - ) - }) - }); - parser.finish(); - - let mut pipeline = c.benchmark_group("engine_stats_stream_pipeline"); - pipeline.sample_size(10); - pipeline.throughput(Throughput::Elements(event_count)); - pipeline.bench_function(collector_benchmark_name, |b| { - b.iter_custom(|iters| { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_time() - .build() - .expect("benchmark runtime should build"); - let mut total = Duration::ZERO; - for _ in 0..iters { - total += runtime.block_on(ingest_and_apply_request_counters(event_count)); - } - total - }) - }); - pipeline.bench_function(endpoint_benchmark_name, |b| { - let events = Arc::new(endpoint_event_lines(event_count)); - b.iter_custom(|iters| { - let runtime = tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .expect("benchmark runtime should build"); - let mut total = Duration::ZERO; - for _ in 0..iters { - total += runtime.block_on(ingest_endpoint_to_collector(events.clone())); - } - total - }) - }); - pipeline.finish(); -} - -fn running_in_criterion_test_mode() -> bool { - let arguments: Vec<_> = std::env::args_os().collect(); - // Criterion treats invocations without `--bench` as cargo-test smoke runs. - !arguments.iter().any(|arg| arg == "--bench") || arguments.iter().any(|arg| arg == "--test") -} - -async fn ingest_and_apply_request_counters(event_count: u64) -> Duration { - let (runtime_state, stats_update_tx, collector) = start_benchmark_collector(); - - let observed_start = TokioInstant::now(); - let started_at = Instant::now(); - for index in 0..event_count { - let request_index = index % REQUEST_IDS; - let step = index / REQUEST_IDS + 1; - stats_update_tx - .send_async(request_update( - format!("req-{request_index}"), - step * 8, - step, - false, - observed_start + Duration::from_millis(index), - )) - .await - .expect("stats collector should receive benchmark update"); - } - send_sentinel_updates(&stats_update_tx, observed_start, event_count).await; - wait_for_sentinel_snapshot(&runtime_state).await; - let elapsed = started_at.elapsed(); - - collector.shutdown().await; - elapsed -} - -async fn ingest_endpoint_to_collector(events: Arc>) -> Duration { - let (runtime_state, stats_update_tx, collector) = start_benchmark_collector(); - let (base_url, endpoint) = start_stats_endpoint(events.clone()).await; - let mut stream_config = EngineStatsStreamConfig::new( - &base_url, - "/pylon/v1/stats/stream", - EngineStatsStreamMode::Required, - ); - stream_config.initial_reconnect_backoff = Duration::from_secs(60); - stream_config.max_reconnect_backoff = Duration::from_secs(60); - let stream = start_engine_stats_stream(stream_config, stats_update_tx) - .expect("benchmark stats stream should start"); - - let started_at = Instant::now(); - wait_for_sentinel_snapshot(&runtime_state).await; - let elapsed = started_at.elapsed(); - - stream.shutdown().await; - collector.shutdown().await; - endpoint.abort(); - let _ = endpoint.await; - elapsed -} - -fn start_benchmark_collector() -> ( - PylonRuntimeState, - flume::Sender, - StatsCollectorHandle, -) { - let config = StatsCollectorConfig { - observation_channel_capacity: 4_096, - engine_stats_request_ttl: Duration::from_secs(300), - engine_stats_model_ttl: Duration::from_secs(300), - ..Default::default() - }; - let (runtime_state, observation_rx) = PylonRuntimeState::observed( - stargate_proto::pb::InferenceServerStatus::Unknown, - &[], - config.observation_channel_capacity, - None, - ); - let (stats_update_tx, stats_update_rx) = stats_aggregator_update_channel(&config); - let collector = start_stats_collector_with_engine_stats( - config, - observation_rx, - Some(stats_update_rx), - runtime_state.clone(), - ); - (runtime_state, stats_update_tx, collector) -} - -async fn start_stats_endpoint(events: Arc>) -> (String, JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("benchmark stats endpoint should bind"); - let addr = listener - .local_addr() - .expect("benchmark stats endpoint should have local addr"); - let app = Router::new() - .route("/pylon/v1/stats/stream", get(stats_endpoint)) - .with_state(events); - let handle = tokio::spawn(async move { - axum::serve(listener, app) - .await - .expect("benchmark stats endpoint should serve"); - }); - (format!("http://{addr}"), handle) -} - -async fn stats_endpoint(State(events): State>>) -> Response { - let stream = stream::iter(0..events.len()).map(move |index| { - let event = events[index].clone(); - Ok::(event) - }); - Response::builder() - .header("content-type", "application/x-ndjson") - .body(Body::from_stream(stream)) - .expect("benchmark stats endpoint response should build") -} - -fn endpoint_event_lines(event_count: u64) -> Vec { - let mut events = Vec::with_capacity(usize::try_from(event_count).unwrap_or_default() + 2); - events.push(Bytes::from(stats_event_line_for( - "req-sentinel", - 0, - 0, - false, - ))); - events.extend((0..event_count).map(|index| Bytes::from(stats_event_line(index)))); - events.push(Bytes::from(stats_event_line_for( - "req-sentinel", - 0, - SENTINEL_OUTPUT_TOKENS, - true, - ))); - events -} - -fn stats_event_line(index: u64) -> String { - let request_index = index % REQUEST_IDS; - let step = index / REQUEST_IDS + 1; - format!( - "{{\"v\":1,\"type\":\"stats\",\"request_id\":\"req-{request_index}\",\"model\":\"model-a\",\"tokens_processed\":{},\"tokens_generated\":{}}}\n", - step * 8, - step - ) -} - -fn stats_event_line_for( - request_id: &str, - tokens_processed: u64, - tokens_generated: u64, - finished: bool, -) -> String { - format!( - "{{\"v\":1,\"type\":\"stats\",\"request_id\":\"{request_id}\",\"model\":\"model-a\",\"tokens_processed\":{tokens_processed},\"tokens_generated\":{tokens_generated},\"finished\":{finished}}}\n", - ) -} - -async fn send_sentinel_updates( - stats_update_tx: &flume::Sender, - observed_start: TokioInstant, - event_count: u64, -) { - let sentinel_start = observed_start + Duration::from_millis(event_count); - stats_update_tx - .send_async(request_update("req-sentinel", 0, 0, false, sentinel_start)) - .await - .expect("stats collector should receive sentinel start"); - stats_update_tx - .send_async(request_update( - "req-sentinel", - 0, - SENTINEL_OUTPUT_TOKENS, - true, - sentinel_start + Duration::from_secs(1), - )) - .await - .expect("stats collector should receive sentinel finish"); -} - -fn request_update( - request_id: impl Into, - tokens_processed: u64, - tokens_generated: u64, - finished: bool, - observed_at: TokioInstant, -) -> StatsAggregatorUpdate { - StatsAggregatorUpdate::RequestCounters(RequestCounterUpdate::new(RequestCounterUpdateInput { - source: StatsUpdateSource::EngineStatsStream, - request_id: request_id.into(), - model_id: "model-a".to_string(), - tokens_processed: Some(tokens_processed), - tokens_generated: Some(tokens_generated), - finished, - observed_at, - })) -} - -async fn wait_for_sentinel_snapshot(runtime_state: &PylonRuntimeState) { - let receive_sentinel = async { - let mut poll = tokio::time::interval(Duration::from_millis(1)); - loop { - poll.tick().await; - if runtime_state - .model_stats("model-a") - .is_some_and(|stats| stats.max_output_tps >= SENTINEL_OUTPUT_TPS) - { - break; - } - } - }; - if running_in_criterion_test_mode() { - tokio::time::timeout(TEST_SENTINEL_TIMEOUT, receive_sentinel) - .await - .expect("sentinel stats snapshot should be published in benchmark smoke mode"); - } else { - receive_sentinel.await; - } -} - -criterion_group!(benches, bench_engine_stats_stream); -criterion_main!(benches); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs index 2365f7153..7b2d5e14f 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/lib.rs @@ -63,9 +63,8 @@ pub use runtime_state::{ pub use stats::{ EngineStatsStreamConfig, EngineStatsStreamHandle, EngineStatsStreamMode, MetricsServerHandle, PylonMetrics, RequestCounterUpdate, RequestCounterUpdateInput, StatsAggregatorUpdate, - StatsCollectorConfig, StatsCollectorHandle, StatsUpdateSource, - parse_engine_stats_line_for_benchmark, start_engine_stats_stream, start_metrics_server, - start_stats_collector, start_stats_collector_with_engine_stats, + StatsCollectorConfig, StatsCollectorHandle, StatsUpdateSource, start_engine_stats_stream, + start_metrics_server, start_stats_collector, start_stats_collector_with_engine_stats, stats_aggregator_update_channel, }; pub use upstream_health::{DEFAULT_UPSTREAM_HEALTH_PATHS, UpstreamHealthPaths}; diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs index cd7ceedf9..0daf85ff6 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs @@ -848,73 +848,6 @@ mod tests { } } - #[derive(Clone)] - struct BlockingKvCacheState { - blocked: Arc, - blocked_polls: Arc, - released: Arc, - } - - struct BlockingKvCacheServer { - server: TestHttpServer, - blocked: Arc, - blocked_polls: Arc, - released: Arc, - } - - impl BlockingKvCacheServer { - async fn spawn() -> Self { - let blocked = Arc::new(AtomicBool::new(false)); - let blocked_polls = Arc::new(AtomicUsize::new(0)); - let released = Arc::new(Notify::new()); - let state = BlockingKvCacheState { - blocked: blocked.clone(), - blocked_polls: blocked_polls.clone(), - released: released.clone(), - }; - let server = TestHttpServer::spawn( - Router::new() - .route("/kv-cache", get(test_kv_cache)) - .with_state(state), - ) - .await; - Self { - server, - blocked, - blocked_polls, - released, - } - } - - fn url(&self) -> String { - format!("{}/kv-cache", self.server.as_str()) - } - - fn block(&self) { - self.blocked.store(true, Ordering::SeqCst); - } - - fn unblock(&self) { - self.blocked.store(false, Ordering::SeqCst); - self.released.notify_waiters(); - } - - fn blocked_poll_count(&self) -> usize { - self.blocked_polls.load(Ordering::SeqCst) - } - - async fn wait_for_blocked_poll_after(&self, count: usize) { - wait_for("KV-cache stats poll should block", || { - self.blocked_polls.load(Ordering::SeqCst) > count - }) - .await; - } - - async fn shutdown(self) { - self.server.shutdown().await; - } - } - async fn test_models(State(state): State) -> Response { state.discovery_polls.fetch_add(1, Ordering::SeqCst); state.discovery_polled.notify_one(); @@ -986,26 +919,6 @@ mod tests { } } - async fn test_kv_cache(State(state): State) -> Response { - if state.blocked.load(Ordering::SeqCst) { - state.blocked_polls.fetch_add(1, Ordering::SeqCst); - loop { - let released = state.released.notified(); - if !state.blocked.load(Ordering::SeqCst) { - break; - } - released.await; - } - } - Json(json!({ - "model": "model-a", - "kv_cache_capacity_tokens": 1000, - "kv_cache_used_tokens": 400, - "kv_cache_free_tokens": 600 - })) - .into_response() - } - fn discovery_config(base_url: &str) -> ModelLifecycleConfig { ModelLifecycleConfig { upstream_http_base_url: base_url.to_string(), @@ -1272,15 +1185,7 @@ mod tests { #[tokio::test] async fn retire_cancels_canary_before_removing_runtime_generation() { - let kv_cache = BlockingKvCacheServer::spawn().await; - let stats_config = StatsCollectorConfig { - kv_cache_stats_url: Some(kv_cache.url()), - // Reconnect quickly and keep the request open so retirement parks on - // stats cleanup after the canary ordering point under test. - kv_cache_reconnect_interval: Duration::from_millis(1), - kv_cache_connect_timeout: Duration::from_secs(60), - ..StatsCollectorConfig::default() - }; + let stats_config = StatsCollectorConfig::default(); let (runtime_state, observations) = PylonRuntimeState::observed( InferenceServerStatus::Active, &[], @@ -1344,9 +1249,6 @@ mod tests { next_poll_at: None, }; - let blocked_polls = kv_cache.blocked_poll_count(); - kv_cache.block(); - kv_cache.wait_for_blocked_poll_after(blocked_polls).await; let retire_generation = generation.clone(); let retire = tokio::spawn(async move { supervisor.retire(&retire_generation).await; @@ -1367,12 +1269,8 @@ mod tests { "normal retirement must not look like an unexpected canary exit" ); - kv_cache.unblock(); - let _supervisor = retire - .await - .expect("retirement task should not panic after stats resumes"); + let _supervisor = retire.await.expect("retirement task should not panic"); stats.shutdown().await; - kv_cache.shutdown().await; } #[tokio::test] diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs index e24285121..75e67d7ae 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs @@ -58,19 +58,19 @@ pub(super) struct GenerationMetricsState { pub(super) pinned_input_tps: Option, } #[derive(Debug, Clone, Default, PartialEq, Eq)] -pub(super) struct KvCacheStatsSnapshot { - pub(super) model: String, - pub(super) aliases: Vec, - pub(super) kv_cache_capacity_tokens: u64, - pub(super) kv_cache_used_tokens: u64, - pub(super) kv_cache_free_tokens: u64, - pub(super) source_observed_at_unix_ms: u64, - pub(super) complete: bool, +pub struct KvCacheStatsSnapshot { + pub(crate) model: String, + pub(crate) aliases: Vec, + pub(crate) kv_cache_capacity_tokens: u64, + pub(crate) kv_cache_used_tokens: u64, + pub(crate) kv_cache_free_tokens: u64, + pub(crate) source_observed_at_unix_ms: u64, + pub(crate) complete: bool, } #[derive(Debug, Clone, Default, PartialEq, Eq)] -pub(super) struct KvCacheStatsEnvelope { - pub(super) models: Vec, +pub struct KvCacheStatsEnvelope { + pub(crate) models: Vec, } struct RequestCounterState { generation: ModelGeneration, @@ -331,6 +331,8 @@ impl StatsAggregator { StatsAggregatorUpdate::RequestCounters(update) => { self.apply_request_counters_into(update, updated_models) } + StatsAggregatorUpdate::KvCache(snapshot) => updated_models + .extend(self.apply_kv_cache_snapshot(snapshot, tokio::time::Instant::now())), StatsAggregatorUpdate::FinalizeRequest(update) => { updated_models.extend(self.finalize_request(update)) } diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs index 650c4fad8..8520c8c70 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs @@ -24,16 +24,13 @@ use crate::runtime_state::ModelGeneration; use crate::{CurrentModelStats, PylonRuntimeState, RequestObservationEvent}; use stargate_runtime::OwnedTask; -use super::aggregator::{ENGINE_STATS_SOURCE, StatsAggregator}; -use super::kv_stats_stream::{KvStatsStreamConfig, run_kv_stats_stream}; +use super::aggregator::{ENGINE_STATS_SOURCE, KvCacheStatsEnvelope, StatsAggregator}; const DEFAULT_OBSERVATION_CHANNEL_CAPACITY: usize = 1024; const DEFAULT_SMOOTHING_WINDOW_SIZE: usize = 8; const DEFAULT_MIN_INPUT_TOKENS: u64 = 1; const DEFAULT_MIN_OUTPUT_TOKENS: u64 = 1; const DEFAULT_DURATION_FLOOR: Duration = Duration::from_millis(10); -const DEFAULT_KV_CACHE_RECONNECT_INTERVAL: Duration = Duration::from_secs(1); -const DEFAULT_KV_CACHE_CONNECT_TIMEOUT: Duration = Duration::from_secs(1); const DEFAULT_KV_CACHE_STATS_TTL: Duration = Duration::from_secs(5); const DEFAULT_ENGINE_STATS_REQUEST_TTL: Duration = Duration::from_secs(300); const DEFAULT_ENGINE_STATS_MODEL_TTL: Duration = Duration::from_secs(30); @@ -46,9 +43,6 @@ pub struct StatsCollectorConfig { pub min_input_tokens: u64, pub min_output_tokens: u64, pub duration_floor: Duration, - pub kv_cache_stats_url: Option, - pub kv_cache_reconnect_interval: Duration, - pub kv_cache_connect_timeout: Duration, pub kv_cache_stats_ttl: Duration, pub engine_stats_request_ttl: Duration, pub engine_stats_model_ttl: Duration, @@ -64,9 +58,6 @@ impl Default for StatsCollectorConfig { min_input_tokens: DEFAULT_MIN_INPUT_TOKENS, min_output_tokens: DEFAULT_MIN_OUTPUT_TOKENS, duration_floor: DEFAULT_DURATION_FLOOR, - kv_cache_stats_url: None, - kv_cache_reconnect_interval: DEFAULT_KV_CACHE_RECONNECT_INTERVAL, - kv_cache_connect_timeout: DEFAULT_KV_CACHE_CONNECT_TIMEOUT, kv_cache_stats_ttl: DEFAULT_KV_CACHE_STATS_TTL, engine_stats_request_ttl: DEFAULT_ENGINE_STATS_REQUEST_TTL, engine_stats_model_ttl: DEFAULT_ENGINE_STATS_MODEL_TTL, @@ -199,6 +190,7 @@ pub enum StatsUpdateSource { #[derive(Debug, Clone)] pub enum StatsAggregatorUpdate { RequestCounters(RequestCounterUpdate), + KvCache(KvCacheStatsEnvelope), FinalizeRequest(FinalizeRequestUpdate), EnableOpenAiFallback, } @@ -333,31 +325,6 @@ async fn run_stats_collector( ) { let config = aggregator.config.clone(); let runtime_state = aggregator.runtime_state.clone(); - let kv_stream_stop = stop.child_token(); - let (mut kv_stats_rx, kv_stream_task) = if let Some(url) = config.kv_cache_stats_url.clone() { - let (tx, rx) = flume::bounded(config.observation_channel_capacity); - let stream_config = KvStatsStreamConfig { - url, - reconnect_interval: config.kv_cache_reconnect_interval, - connect_timeout: config.kv_cache_connect_timeout, - idle_timeout: if config.kv_cache_stats_ttl.is_zero() { - DEFAULT_KV_CACHE_STATS_TTL - } else { - config.kv_cache_stats_ttl - }, - }; - let stream_stop = kv_stream_stop.clone(); - ( - Some(rx), - Some(tokio::spawn(run_kv_stats_stream( - stream_config, - tx, - stream_stop, - ))), - ) - } else { - (None, None) - }; let mut engine_stats_sweep = tokio::time::interval(config.engine_stats_sweep_interval); engine_stats_sweep.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); let mut stats_aggregator_updated_models = Vec::with_capacity(2); @@ -390,19 +357,6 @@ async fn run_stats_collector( } apply_collector_command(&mut aggregator, &runtime_state, command); } - snapshot = async { - match &kv_stats_rx { - Some(rx) => rx.recv_async().await.ok(), - None => std::future::pending().await, - } - } => { - let Some(snapshot) = snapshot else { - kv_stats_rx = None; - continue; - }; - let updates = aggregator.apply_kv_cache_snapshot(snapshot, TokioInstant::now()); - publish_model_stats_updates(&runtime_state, updates); - } event = observation_rx.recv_async() => { let Ok(event) = event else { break 'collector; @@ -459,11 +413,6 @@ async fn run_stats_collector( } } } - - kv_stream_stop.cancel(); - if let Some(task) = kv_stream_task { - let _ = task.await; - } } fn publish_observation_event( @@ -555,7 +504,6 @@ fn retain_latest_model_updates( mod tests { use std::collections::HashMap; use std::sync::Arc; - use std::sync::atomic::{AtomicUsize, Ordering}; use super::super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot, StatsAggregator}; use super::super::metrics::PylonMetrics; @@ -565,8 +513,6 @@ mod tests { use crate::generated_request_id::{GeneratedRequestKind, next_generated_request_id}; use crate::request_observer::RequestObservationEndpoint; use crate::request_observer::RequestObservationState; - use axum::{Json, Router, routing::get}; - use tokio::net::TcpListener; const MODEL_STATS_TEST_TIMEOUT: Duration = milliseconds(500); @@ -1219,21 +1165,6 @@ mod tests { assert!(body.contains(expected), "{context}"); } - async fn spawn_kv_cache_server( - app: Router, - ) -> (std::net::SocketAddr, tokio::task::JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("listener should bind"); - let address = listener.local_addr().expect("listener should have address"); - let server = tokio::spawn(async move { - axum::serve(listener, app) - .await - .expect("KV-cache test server should run"); - }); - (address, server) - } - #[test] fn stats_stream_cumulative_request_counters_drive_stats_aggregator() { let mut aggregator = test_aggregator(StatsCollectorConfig::default()); @@ -2376,25 +2307,15 @@ mod tests { } #[tokio::test] - async fn kv_cache_stream_updates_model_metrics() { - const SNAPSHOT: &str = "{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":42,\"models\":[{\"model\":\"model-a\",\"aliases\":[],\"routing_cache\":{\"role\":\"decode\",\"capacity_tokens\":1000,\"used_tokens\":400,\"free_tokens\":600},\"pools\":[]}]}\n"; - let requests = Arc::new(AtomicUsize::new(0)); - let handler_requests = requests.clone(); + async fn mixed_stream_kv_update_updates_model_metrics() { let metrics = PylonMetrics::new().expect("metrics should initialize"); - let app = Router::new().route( - "/kv-cache", - get(move || { - handler_requests.fetch_add(1, Ordering::Relaxed); - async { SNAPSHOT } - }), - ); - let (addr, server) = spawn_kv_cache_server(app).await; - let config = config!( - kv_cache_stats_url: Some(format!("http://{addr}/kv-cache")), - kv_cache_reconnect_interval: milliseconds(10), - kv_cache_connect_timeout: seconds(1), - ); - let collector = RunningCollector::spawn(config, Some(metrics.clone()), false); + let collector = + RunningCollector::spawn(StatsCollectorConfig::default(), Some(metrics.clone()), true); + let mut snapshot = kv_cache_stats("model-a"); + snapshot.source_observed_at_unix_ms = 42; + collector + .send_update(StatsAggregatorUpdate::KvCache(kv_cache_envelope(snapshot))) + .await; let stats = collector .wait_for_stats("KV-cache stats should be published", |stats| { stats.kv_cache_capacity_tokens == 1000 @@ -2412,44 +2333,9 @@ mod tests { assert!(body.contains(r#"pylon_model_kv_cache_capacity_tokens{model="model-a"} 1000"#)); assert!(body.contains(r#"pylon_model_kv_cache_used_tokens{model="model-a"} 400"#)); assert!(body.contains(r#"pylon_model_kv_cache_free_tokens{model="model-a"} 600"#)); - tokio::time::timeout(seconds(1), async { - while requests.load(Ordering::Relaxed) < 2 { - tokio::time::sleep(milliseconds(10)).await; - } - }) - .await - .expect("finite KV stats responses should reconnect"); tokio::time::timeout(seconds(2), collector.handle.shutdown()) .await .expect("collector should stop"); - server.abort(); - } - - #[tokio::test] - async fn stats_collector_shutdown_interrupts_blocked_kv_cache_connect() { - let poll_entered = Arc::new(tokio::sync::Barrier::new(2)); - let server_poll_entered = poll_entered.clone(); - let app = Router::new().route( - "/kv-cache", - get(move || { - let poll_entered = server_poll_entered.clone(); - async move { - poll_entered.wait().await; - std::future::pending::>().await - } - }), - ); - let (addr, server) = spawn_kv_cache_server(app).await; - let config = config!( - kv_cache_stats_url: Some(format!("http://{addr}/kv-cache")), - kv_cache_reconnect_interval: milliseconds(1), - kv_cache_connect_timeout: seconds(60), - ); - let collector = RunningCollector::spawn(config, None, false); - poll_entered.wait().await; - let stopped = tokio::time::timeout(seconds(1), collector.handle.shutdown()).await; - server.abort(); - stopped.expect("collector shutdown should interrupt blocked KV-cache connect"); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs index 607f68640..584f030ac 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs @@ -1,49 +1,33 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. -use std::borrow::Cow; use std::fmt; use std::str::FromStr; use std::sync::Arc; use std::time::Duration; -use bytes::Bytes; -use futures::{Stream, StreamExt}; -use reqwest::StatusCode; -use serde::de::{IgnoredAny, MapAccess, SeqAccess, Visitor}; -use serde::{Deserialize, Deserializer}; +use stargate_proto::dynamo_frontend_stats as proto; +use stargate_runtime::OwnedTask; use tokio::time::Instant as TokioInstant; use tokio_util::sync::CancellationToken; +use tonic::Code; use super::collector::{RequestCounterUpdate, StatsAggregatorUpdate, StatsUpdateSource}; +use super::kv_stats::kv_snapshot_from_proto; use super::metrics::PylonMetrics; use crate::PylonRuntimeState; use crate::generated_request_id::generated_request_generation; -use stargate_runtime::OwnedTask; -const DEFAULT_ENGINE_STATS_STREAM_PATH: &str = "/pylon/v1/stats/stream"; const DEFAULT_INITIAL_RECONNECT_BACKOFF: Duration = Duration::from_millis(100); const DEFAULT_MAX_RECONNECT_BACKOFF: Duration = Duration::from_secs(5); -const DEFAULT_MAX_LINE_BYTES: usize = 64 * 1024; -const HEADER_ACCEPT_NDJSON: &str = "application/x-ndjson"; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum EngineStatsStreamMode { Auto, Required, Off, } + impl EngineStatsStreamMode { pub fn as_str(self) -> &'static str { match self { @@ -53,11 +37,13 @@ impl EngineStatsStreamMode { } } } + impl fmt::Display for EngineStatsStreamMode { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.write_str(self.as_str()) } } + impl FromStr for EngineStatsStreamMode { type Err = ParseEngineStatsStreamModeError; @@ -70,41 +56,40 @@ impl FromStr for EngineStatsStreamMode { } } } -#[derive(Debug, Clone, Copy, thiserror::Error)] + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] #[error("expected one of auto, required, off")] pub struct ParseEngineStatsStreamModeError; + #[derive(Debug, Clone)] pub struct EngineStatsStreamConfig { - pub url: String, + pub endpoint: String, pub mode: EngineStatsStreamMode, pub initial_reconnect_backoff: Duration, pub max_reconnect_backoff: Duration, - pub max_line_bytes: usize, pub metrics: Option>, pub runtime_state: Option, } + impl EngineStatsStreamConfig { - pub fn new(upstream_base_url: &str, path: &str, mode: EngineStatsStreamMode) -> Self { + pub fn new(upstream_base_url: &str, mode: EngineStatsStreamMode) -> Self { Self { - url: crate::upstream_url::upstream_endpoint(upstream_base_url, path), + endpoint: upstream_base_url.trim_end_matches('/').to_string(), mode, initial_reconnect_backoff: DEFAULT_INITIAL_RECONNECT_BACKOFF, max_reconnect_backoff: DEFAULT_MAX_RECONNECT_BACKOFF, - max_line_bytes: DEFAULT_MAX_LINE_BYTES, metrics: None, runtime_state: None, } } } + impl Default for EngineStatsStreamConfig { fn default() -> Self { - Self::new( - "http://127.0.0.1:8090", - DEFAULT_ENGINE_STATS_STREAM_PATH, - EngineStatsStreamMode::Auto, - ) + Self::new("http://127.0.0.1:8090", EngineStatsStreamMode::Auto) } } + owned_task_handle!(EngineStatsStreamHandle); pub fn start_engine_stats_stream( @@ -119,286 +104,46 @@ pub fn start_engine_stats_stream( }); Some(EngineStatsStreamHandle { task }) } -#[derive(Debug)] -pub(crate) enum ParsedEngineStatsEvent { - Stats(RequestCounterUpdate), - Ping, -} -#[derive(Debug, thiserror::Error)] -pub(crate) enum EngineStatsParseError { - #[error("invalid JSON: {0}")] - Json(#[from] serde_json::Error), - #[error("event must be a JSON object")] - NotObject, - #[error("missing field {0}")] - MissingField(&'static str), - #[error("unsupported version {0}")] - UnsupportedVersion(u64), - #[error("invalid field {0}")] - InvalidField(&'static str), - #[error("unknown event type {0}")] - UnknownType(String), - #[error("stats event must include at least one counter unless finished=true")] - EmptyStatsCounters, -} - -pub(crate) fn parse_engine_stats_line( - line: &[u8], - observed_at: TokioInstant, -) -> Result { - if line.iter().find(|byte| !byte.is_ascii_whitespace()) != Some(&b'{') { - serde_json::from_slice::(line)?; - return Err(EngineStatsParseError::NotObject); - } - let mut raw: RawEngineStatsEvent<'_> = serde_json::from_slice(line)?; - let event_type = engine_stats_event_type(&mut raw)?; - match event_type { - "stats" => parse_stats_event(raw, observed_at), - "ping" => Ok(ParsedEngineStatsEvent::Ping), - other => Err(EngineStatsParseError::UnknownType(other.to_string())), - } -} -fn engine_stats_event_type<'a>( - raw: &'a mut RawEngineStatsEvent<'_>, -) -> Result<&'a str, EngineStatsParseError> { - let version = - optional_u64(raw.version.take(), "v")?.ok_or(EngineStatsParseError::MissingField("v"))?; - if version != 1 { - return Err(EngineStatsParseError::UnsupportedVersion(version)); - } - match raw.event_type.as_ref() { - Some(JsonScalar::String(value)) => Ok(value.as_ref()), - Some(_) => Err(EngineStatsParseError::InvalidField("type")), - None => Err(EngineStatsParseError::MissingField("type")), - } -} -fn parse_stats_event( - raw: RawEngineStatsEvent<'_>, - observed_at: TokioInstant, -) -> Result { - let request_id = required_nonempty_string(raw.request_id, "request_id")?; - let model_id = required_nonempty_string(raw.model, "model")?; - let tokens_processed = optional_u64(raw.tokens_processed, "tokens_processed")?; - let tokens_generated = optional_u64(raw.tokens_generated, "tokens_generated")?; - let finished = match raw.finished { - Some(JsonScalar::Bool(value)) => value, - Some(_) => return Err(EngineStatsParseError::InvalidField("finished")), - None => false, - }; - if tokens_processed.is_none() && tokens_generated.is_none() && !finished { - return Err(EngineStatsParseError::EmptyStatsCounters); - } - Ok(ParsedEngineStatsEvent::Stats(RequestCounterUpdate { - source: StatsUpdateSource::EngineStatsStream, - request_id, - model_id, - generation: None, - tokens_processed, - tokens_generated, - finished, - observed_at, - })) -} -#[doc(hidden)] -pub fn parse_engine_stats_line_for_benchmark(line: &[u8], observed_at: TokioInstant) -> bool { - parse_engine_stats_line(line, observed_at).is_ok() -} - -#[derive(Default)] -struct RawEngineStatsEvent<'a> { - version: Option>, - event_type: Option>, - request_id: Option>, - model: Option>, - tokens_processed: Option>, - tokens_generated: Option>, - finished: Option>, -} - -impl<'de> Deserialize<'de> for RawEngineStatsEvent<'de> { - fn deserialize(deserializer: D) -> Result - where - D: Deserializer<'de>, - { - deserializer.deserialize_map(RawEngineStatsEventVisitor) - } -} - -struct RawEngineStatsEventVisitor; - -impl<'de> Visitor<'de> for RawEngineStatsEventVisitor { - type Value = RawEngineStatsEvent<'de>; - - fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { - formatter.write_str("an engine stats JSON object") - } - - fn visit_map(self, mut map: M) -> Result - where - M: MapAccess<'de>, - { - let mut event = RawEngineStatsEvent::default(); - while let Some(key) = map.next_key::>()? { - match key.as_ref() { - "v" => event.version = Some(map.next_value()?), - "type" => event.event_type = Some(map.next_value()?), - "request_id" => event.request_id = Some(map.next_value()?), - "model" => event.model = Some(map.next_value()?), - "tokens_processed" => event.tokens_processed = Some(map.next_value()?), - "tokens_generated" => event.tokens_generated = Some(map.next_value()?), - "finished" => event.finished = Some(map.next_value()?), - _ => { - map.next_value::()?; - } - } - } - Ok(event) - } -} - -enum JsonScalar<'a> { - String(Cow<'a, str>), - Unsigned(u64), - Bool(bool), - Invalid, -} - -impl<'de> Deserialize<'de> for JsonScalar<'de> { - fn deserialize(deserializer: D) -> Result - where - D: Deserializer<'de>, - { - deserializer.deserialize_any(JsonScalarVisitor) - } -} - -struct JsonScalarVisitor; - -impl<'de> Visitor<'de> for JsonScalarVisitor { - type Value = JsonScalar<'de>; - - fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { - formatter.write_str("a JSON scalar") - } - - fn visit_borrowed_str(self, value: &'de str) -> Result { - Ok(JsonScalar::String(Cow::Borrowed(value))) - } - - fn visit_str(self, value: &str) -> Result { - Ok(JsonScalar::String(Cow::Owned(value.to_string()))) - } - - fn visit_u64(self, value: u64) -> Result { - Ok(JsonScalar::Unsigned(value)) - } - - fn visit_i64(self, value: i64) -> Result { - Ok(u64::try_from(value) - .map(JsonScalar::Unsigned) - .unwrap_or(JsonScalar::Invalid)) - } - - fn visit_bool(self, value: bool) -> Result { - Ok(JsonScalar::Bool(value)) - } - - fn visit_seq(self, mut seq: A) -> Result - where - A: SeqAccess<'de>, - { - while seq.next_element::()?.is_some() {} - Ok(JsonScalar::Invalid) - } - - fn visit_map(self, mut map: M) -> Result - where - M: MapAccess<'de>, - { - while map.next_entry::()?.is_some() {} - Ok(JsonScalar::Invalid) - } - - fn visit_f64(self, _value: f64) -> Result { - Ok(JsonScalar::Invalid) - } - - fn visit_unit(self) -> Result { - Ok(JsonScalar::Invalid) - } -} - -fn optional_u64( - value: Option>, - field: &'static str, -) -> Result, EngineStatsParseError> { - match value { - Some(JsonScalar::Unsigned(value)) => Ok(Some(value)), - Some(_) => Err(EngineStatsParseError::InvalidField(field)), - None => Ok(None), - } -} - -fn required_nonempty_string( - value: Option>, - field: &'static str, -) -> Result { - let value = match value { - Some(JsonScalar::String(value)) => value, - Some(_) => return Err(EngineStatsParseError::InvalidField(field)), - None => return Err(EngineStatsParseError::MissingField(field)), - }; - let value = value.trim(); - if value.is_empty() { - return Err(EngineStatsParseError::InvalidField(field)); - } - Ok(value.to_string()) -} async fn run_engine_stats_stream( config: EngineStatsStreamConfig, stats_update_tx: flume::Sender, stop: CancellationToken, ) { - let client = reqwest::Client::new(); let mut backoff = config.initial_reconnect_backoff; let mut valid_event_seen = false; loop { if stop.is_cancelled() { return; } - let retry_reason = match read_stream_once( - &config, - &client, - &stats_update_tx, - &stop, - &mut valid_event_seen, - ) - .await - { - StreamReadOutcome::Stopped => return, - StreamReadOutcome::Unsupported - if config.mode == EngineStatsStreamMode::Auto && !valid_event_seen => - { - tracing::warn!( - url = config.url, - "engine stats stream unsupported; using OpenAI fallback observation" - ); - let _ = send_stats_update( - &stats_update_tx, - StatsAggregatorUpdate::EnableOpenAiFallback, - &stop, - ) - .await; - return; - } - StreamReadOutcome::Unsupported => "unsupported", - StreamReadOutcome::Retry(reason) => reason, - }; + let retry_reason = + match read_stream_once(&config, &stats_update_tx, &stop, &mut valid_event_seen).await { + StreamReadOutcome::Stopped => return, + StreamReadOutcome::Unsupported + if config.mode == EngineStatsStreamMode::Auto && !valid_event_seen => + { + tracing::warn!( + endpoint = config.endpoint, + "frontend stats gRPC service unsupported; using OpenAI fallback observation" + ); + let _ = send_stats_update( + &stats_update_tx, + StatsAggregatorUpdate::EnableOpenAiFallback, + &stop, + ) + .await; + return; + } + StreamReadOutcome::Unsupported => "unsupported", + StreamReadOutcome::Retry(reason) => reason, + }; + if let Some(metrics) = &config.metrics { metrics.observe_engine_stats_reconnect(retry_reason); } - + if valid_event_seen { + backoff = config.initial_reconnect_backoff; + } if stop .run_until_cancelled(tokio::time::sleep(backoff)) .await @@ -419,214 +164,114 @@ enum StreamReadOutcome { async fn read_stream_once( config: &EngineStatsStreamConfig, - client: &reqwest::Client, stats_update_tx: &flume::Sender, stop: &CancellationToken, valid_event_seen: &mut bool, ) -> StreamReadOutcome { - let response = match open_engine_stats_response(config, client, stop).await { - Ok(response) => response, - Err(outcome) => return outcome, - }; - observe_connected(config, true); - let outcome = - drain_engine_stats_response(config, response, stats_update_tx, stop, valid_event_seen) - .await; - observe_connected(config, false); - outcome -} -async fn drain_engine_stats_response( - config: &EngineStatsStreamConfig, - response: reqwest::Response, - stats_update_tx: &flume::Sender, - stop: &CancellationToken, - valid_event_seen: &mut bool, -) -> StreamReadOutcome { - let mut stream = response.bytes_stream(); - let mut line_buffer = Vec::with_capacity(1024); - let mut discarding_oversized_line = false; - loop { - let chunk = match next_engine_stats_chunk(config, stop, &mut stream).await { - Ok(chunk) => chunk, - Err(outcome) => break outcome, - }; - line_buffer.extend_from_slice(&chunk); - if !process_buffered_engine_stats_lines( - config, - stats_update_tx, - valid_event_seen, - &mut line_buffer, - &mut discarding_oversized_line, - stop, - ) - .await - { - break StreamReadOutcome::Stopped; + let connect = + proto::frontend_stats_client::FrontendStatsClient::connect(config.endpoint.clone()); + let mut client = match stop.run_until_cancelled(connect).await { + None => return StreamReadOutcome::Stopped, + Some(Ok(client)) => client, + Some(Err(error)) => { + tracing::warn!(endpoint = config.endpoint, %error, "frontend stats gRPC connect failed"); + return StreamReadOutcome::Retry("connect_error"); } - } -} -async fn open_engine_stats_response( - config: &EngineStatsStreamConfig, - client: &reqwest::Client, - stop: &CancellationToken, -) -> Result { - let response = tokio::select! { - _ = stop.cancelled() => return Err(StreamReadOutcome::Stopped), - response = client - .get(&config.url) - .header(reqwest::header::ACCEPT, HEADER_ACCEPT_NDJSON) - .send() => response, }; - let response = match response { - Ok(response) => response, - Err(error) => { - tracing::warn!(url = config.url, error = %error, "engine stats stream connect failed"); - return Err(StreamReadOutcome::Retry("connect_error")); - } - }; - if matches!( - response.status(), - StatusCode::NOT_FOUND | StatusCode::METHOD_NOT_ALLOWED | StatusCode::NOT_IMPLEMENTED - ) { - tracing::warn!( - url = config.url, - status = %response.status(), - "engine stats stream endpoint is unsupported" - ); - return Err(StreamReadOutcome::Unsupported); - } - if !response.status().is_success() { - tracing::warn!( - url = config.url, - status = %response.status(), - "engine stats stream returned non-success status" - ); - return Err(StreamReadOutcome::Retry("http_status")); - } - Ok(response) -} -async fn next_engine_stats_chunk( - config: &EngineStatsStreamConfig, - stop: &CancellationToken, - stream: &mut S, -) -> Result -where - S: Stream> + Unpin, -{ - let chunk = tokio::select! { - _ = stop.cancelled() => return Err(StreamReadOutcome::Stopped), - chunk = stream.next() => chunk, + let watch = client.watch_stats(proto::WatchStatsRequest {}); + let response = match stop.run_until_cancelled(watch).await { + None => return StreamReadOutcome::Stopped, + Some(Ok(response)) => response, + Some(Err(status)) => return status_outcome(config, status), }; - let chunk = chunk.ok_or(StreamReadOutcome::Retry("eof"))?; - let chunk = chunk.map_err(|error| { - tracing::warn!(url = config.url, error = %error, "engine stats stream read failed"); - StreamReadOutcome::Retry("read_error") - })?; - (!chunk.is_empty()) - .then_some(chunk) - .ok_or(StreamReadOutcome::Retry("empty_chunk")) -} -async fn process_buffered_engine_stats_lines( - config: &EngineStatsStreamConfig, - stats_update_tx: &flume::Sender, - valid_event_seen: &mut bool, - line_buffer: &mut Vec, - discarding_oversized_line: &mut bool, - stop: &CancellationToken, -) -> bool { - let mut consumed = if *discarding_oversized_line { - let Some(newline_index) = memchr::memchr(b'\n', line_buffer) else { - line_buffer.clear(); - return true; + observe_connected(config, true); + let mut stream = response.into_inner(); + let outcome = loop { + let message = match stop.run_until_cancelled(stream.message()).await { + None => break StreamReadOutcome::Stopped, + Some(message) => message, }; - *discarding_oversized_line = false; - newline_index + 1 - } else { - 0 - }; - while let Some(relative_newline_index) = memchr::memchr(b'\n', &line_buffer[consumed..]) { - let newline_index = consumed + relative_newline_index; - if newline_index - consumed > config.max_line_bytes { - observe_invalid(config, "line_too_large"); - consumed = newline_index + 1; - continue; - } - let line_end = if line_buffer[consumed..newline_index].ends_with(b"\r") { - newline_index - 1 - } else { - newline_index + let update = match message { + Ok(Some(update)) => update, + Ok(None) => break StreamReadOutcome::Retry("eof"), + Err(status) => break status_outcome(config, status), }; - if line_end != consumed { - let event = - parse_engine_stats_line(&line_buffer[consumed..line_end], TokioInstant::now()); - if !emit_engine_stats_event(config, stats_update_tx, valid_event_seen, event, stop) - .await - { - return false; + let (event_type, update) = match translate_update(config, update, TokioInstant::now()) { + Ok(update) => update, + Err(error) => { + tracing::warn!(endpoint = config.endpoint, %error, "invalid frontend stats update"); + if let Some(metrics) = &config.metrics { + metrics.observe_engine_stats_invalid_event("protobuf"); + } + continue; } + }; + *valid_event_seen = true; + if let Some(metrics) = &config.metrics { + metrics.observe_engine_stats_stream_event(event_type); } - consumed = newline_index + 1; - } - *discarding_oversized_line = compact_line_buffer(config, line_buffer, consumed); - true -} - -async fn emit_engine_stats_event( - config: &EngineStatsStreamConfig, - stats_update_tx: &flume::Sender, - valid_event_seen: &mut bool, - event: Result, - stop: &CancellationToken, -) -> bool { - let event = match event { - Ok(event) => event, - Err(error) => { - tracing::warn!(url = config.url, error = %error, "invalid engine stats event"); - observe_invalid(config, error.metric_reason()); - return true; + if !send_stats_update(stats_update_tx, update, stop).await { + break StreamReadOutcome::Stopped; } }; - *valid_event_seen = true; - let (event_type, update) = match event { - ParsedEngineStatsEvent::Stats(update) => ("stats", Some(update)), - ParsedEngineStatsEvent::Ping => ("ping", None), - }; - if let Some(metrics) = &config.metrics { - metrics.observe_engine_stats_stream_event(event_type); - } - match update { - Some(mut update) => { - update.generation = generated_request_generation(&update.request_id, &update.model_id) - .or_else(|| { - config.runtime_state.as_ref().and_then(|runtime_state| { - runtime_state.request_generation(&update.request_id) - }) - }); - send_stats_update( - stats_update_tx, - StatsAggregatorUpdate::RequestCounters(update), - stop, - ) - .await - } - None => true, + observe_connected(config, false); + outcome +} + +fn status_outcome(config: &EngineStatsStreamConfig, status: tonic::Status) -> StreamReadOutcome { + if matches!(status.code(), Code::Unimplemented | Code::NotFound) { + tracing::warn!(endpoint = config.endpoint, %status, "frontend stats gRPC service is unsupported"); + StreamReadOutcome::Unsupported + } else { + tracing::warn!(endpoint = config.endpoint, %status, "frontend stats gRPC stream disconnected"); + StreamReadOutcome::Retry("grpc_status") } } -fn compact_line_buffer( +fn translate_update( config: &EngineStatsStreamConfig, - line_buffer: &mut Vec, - consumed: usize, -) -> bool { - line_buffer.drain(..consumed).count(); - if line_buffer.len() > config.max_line_bytes { - observe_invalid(config, "line_too_large"); - line_buffer.clear(); - true - } else { - false + update: proto::StatsUpdate, + observed_at: TokioInstant, +) -> anyhow::Result<(&'static str, StatsAggregatorUpdate)> { + match update.update { + Some(proto::stats_update::Update::RequestStats(request)) => { + let request_id = request.request_id.trim(); + let model_id = request.model.trim(); + anyhow::ensure!(!request_id.is_empty(), "request ID is empty"); + anyhow::ensure!(!model_id.is_empty(), "model ID is empty"); + anyhow::ensure!( + request.tokens_processed.is_some() + || request.tokens_generated.is_some() + || request.finished, + "request update has no counters" + ); + let generation = generated_request_generation(request_id, model_id).or_else(|| { + config + .runtime_state + .as_ref() + .and_then(|state| state.request_generation(request_id)) + }); + Ok(( + "stats", + StatsAggregatorUpdate::RequestCounters(RequestCounterUpdate { + source: StatsUpdateSource::EngineStatsStream, + request_id: request_id.to_string(), + model_id: model_id.to_string(), + generation, + tokens_processed: request.tokens_processed, + tokens_generated: request.tokens_generated, + finished: request.finished, + observed_at, + }), + )) + } + Some(proto::stats_update::Update::KvStats(snapshot)) => Ok(( + "kv_stats_snapshot", + StatsAggregatorUpdate::KvCache(kv_snapshot_from_proto(snapshot)?), + )), + None => anyhow::bail!("stats update is missing its payload"), } } @@ -645,26 +290,6 @@ async fn send_stats_update( } } -impl EngineStatsParseError { - fn metric_reason(&self) -> &'static str { - match self { - Self::Json(_) => "json", - Self::NotObject => "not_object", - Self::MissingField(_) => "missing_field", - Self::UnsupportedVersion(_) => "version", - Self::InvalidField(_) => "field", - Self::UnknownType(_) => "type", - Self::EmptyStatsCounters => "empty_stats", - } - } -} - -fn observe_invalid(config: &EngineStatsStreamConfig, reason: &'static str) { - if let Some(metrics) = &config.metrics { - metrics.observe_engine_stats_invalid_event(reason); - } -} - fn observe_connected(config: &EngineStatsStreamConfig, connected: bool) { if let Some(metrics) = &config.metrics { metrics.observe_engine_stats_stream_connected(config.mode.as_str(), connected); @@ -674,640 +299,201 @@ fn observe_connected(config: &EngineStatsStreamConfig, connected: bool) { #[cfg(test)] mod tests { use super::*; - use axum::{Router, routing::get}; - use std::sync::{ - Arc, - atomic::{AtomicUsize, Ordering}, - }; - use tokio::net::TcpListener; - - use crate::generated_request_id::{GeneratedRequestKind, next_generated_request_id}; - use crate::request_observer::{ - RequestObservationEndpoint, RequiredTunnelHeaders, TunnelRequestObserver, - }; - - fn parse(line: &[u8]) -> Result { - parse_engine_stats_line(line, TokioInstant::now()) - } + use std::pin::Pin; + use std::sync::atomic::{AtomicUsize, Ordering}; - struct ProcessedLines { - valid_event_seen: bool, - updates: flume::Receiver, - } - - impl ProcessedLines { - fn assert_state(&self, valid_event_seen: bool) { - assert_eq!(self.valid_event_seen, valid_event_seen); - } - - fn assert_counter(self, context: &str) { - self.assert_state(true); - assert!(matches!( - self.updates - .try_recv() - .unwrap_or_else(|_| panic!("{context}")), - StatsAggregatorUpdate::RequestCounters(_) - )); - } - - fn assert_no_update(self, valid_event_seen: bool) { - self.assert_state(valid_event_seen); - assert!(self.updates.try_recv().is_err()); - } - } - - async fn process_lines( - config: &EngineStatsStreamConfig, - chunks: impl IntoIterator>, - ) -> ProcessedLines { - let (tx, updates) = flume::bounded(4); - let mut valid_event_seen = false; - let mut remaining = Vec::new(); - let mut discarding_oversized_line = false; - for chunk in chunks { - remaining.extend_from_slice(chunk.as_ref()); - assert!( - process_buffered_engine_stats_lines( - config, - &tx, - &mut valid_event_seen, - &mut remaining, - &mut discarding_oversized_line, - &CancellationToken::new(), - ) - .await - ); - } - assert!(remaining.is_empty()); - assert!(!discarding_oversized_line); - ProcessedLines { - valid_event_seen, - updates, - } - } - - async fn serve_stats(app: Router) -> (String, tokio::task::JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("listener should bind"); - let base_url = format!( - "http://{}", - listener.local_addr().expect("listener should have addr") - ); - let server = tokio::spawn(async move { - axum::serve(listener, app).await.expect("server should run"); - }); - (base_url, server) - } + use futures::Stream; + use tokio::net::TcpListener; - fn reconnecting_config(base_url: &str, mode: EngineStatsStreamMode) -> EngineStatsStreamConfig { - EngineStatsStreamConfig { - initial_reconnect_backoff: Duration::from_millis(1), - max_reconnect_backoff: Duration::from_millis(1), - ..EngineStatsStreamConfig::new(base_url, "/pylon/v1/stats/stream", mode) + fn request_update(request: proto::RequestStats) -> proto::StatsUpdate { + proto::StatsUpdate { + update: Some(proto::stats_update::Update::RequestStats(request)), } } - async fn receive_update( - updates: &flume::Receiver, - context: &str, - ) -> StatsAggregatorUpdate { - tokio::time::timeout(Duration::from_secs(2), updates.recv_async()) - .await - .unwrap_or_else(|_| panic!("{context}")) - .expect("stats update should be sent") - } - #[test] - fn parses_valid_engine_stats_events() { - let event = parse( - br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","tokens_processed":4096,"tokens_generated":128,"finished":true}"#, - ) - .expect("stats event should parse"); - let ParsedEngineStatsEvent::Stats(update) = event else { - panic!("expected request counters update"); - }; - assert_eq!(update.request_id, "req-1"); - assert_eq!(update.model_id, "llama"); - assert_eq!(update.tokens_processed, Some(4096)); - assert_eq!(update.tokens_generated, Some(128)); - assert!(update.finished); - - assert!(matches!( - parse(br#"{"v":1,"type":"ping"}"#).expect("ping should parse"), - ParsedEngineStatsEvent::Ping - )); - } - - #[tokio::test] - async fn calibration_request_ids_route_engine_events_to_their_exact_generation() { - let runtime_state = PylonRuntimeState::new( - stargate_proto::pb::InferenceServerStatus::Active, - &["model-a".to_string()], - ); - let generation = runtime_state - .current_generation("model-a") - .expect("test generation should exist"); - let request_id = next_generated_request_id(GeneratedRequestKind::Calibration, &generation); - let config = EngineStatsStreamConfig { - runtime_state: Some(runtime_state), - ..EngineStatsStreamConfig::default() - }; - let line = format!( - "{{\"v\":1,\"type\":\"stats\",\"request_id\":\"{request_id}\",\"model\":\"model-a\",\"tokens_processed\":64}}\n" - ); - - let processed = process_lines(&config, [line]).await; - let StatsAggregatorUpdate::RequestCounters(update) = processed - .updates - .try_recv() - .expect("calibration event should enter the ordinary stats pipeline") - else { - panic!("expected request counters update"); - }; - - assert_eq!(update.generation, Some(generation)); + fn parses_stream_modes() { + assert_eq!("auto".parse(), Ok(EngineStatsStreamMode::Auto)); + assert_eq!("required".parse(), Ok(EngineStatsStreamMode::Required)); + assert_eq!("off".parse(), Ok(EngineStatsStreamMode::Off)); + assert!("other".parse::().is_err()); } - #[tokio::test] - async fn platform_request_id_routes_engine_stats_to_the_live_generation() { - let runtime_state = PylonRuntimeState::new( - stargate_proto::pb::InferenceServerStatus::Active, - &["platform-model".to_string()], - ); - let generation = runtime_state - .current_generation("platform-model") - .expect("test generation should exist"); - let observer = TunnelRequestObserver::accepted( - RequestObservationEndpoint::ChatCompletions, - RequiredTunnelHeaders { - request_id: "gateway-request".to_string(), - routing_key: None, - model_id: "platform-model".to_string(), - priority: None, - input_tokens: 64, - accepted_at: std::time::Instant::now(), - }, - Some(generation.clone()), - runtime_state.clone(), - ); - let config = EngineStatsStreamConfig { - runtime_state: Some(runtime_state), - ..EngineStatsStreamConfig::default() - }; - - let processed = process_lines( - &config, - [r#"{"v":1,"type":"stats","request_id":"gateway-request","model":"platform-model","tokens_processed":64,"finished":true} -"#], + #[test] + fn translates_request_counters() { + let update = request_update(proto::RequestStats { + request_id: "request-a".to_string(), + model: "model-a".to_string(), + tokens_processed: Some(12), + tokens_generated: Some(3), + finished: false, + }); + let (_, update) = translate_update( + &EngineStatsStreamConfig::default(), + update, + TokioInstant::now(), ) - .await; - let StatsAggregatorUpdate::RequestCounters(update) = processed - .updates - .try_recv() - .expect("engine event should enter the stats pipeline") - else { - panic!("expected request counters update"); + .unwrap(); + let StatsAggregatorUpdate::RequestCounters(update) = update else { + panic!("expected request counters"); }; - - assert_eq!(update.request_id, "gateway-request"); - assert_eq!(update.model_id, "platform-model"); - assert_eq!(update.generation, Some(generation)); - drop(observer); + assert_eq!(update.request_id, "request-a"); + assert_eq!(update.model_id, "model-a"); + assert_eq!(update.tokens_processed, Some(12)); + assert_eq!(update.tokens_generated, Some(3)); } #[test] - fn rejects_invalid_engine_stats_events() { - assert!(matches!( - parse(br#"{"v":2,"type":"ping"}"#).unwrap_err(), - EngineStatsParseError::UnsupportedVersion(2) - )); - assert!(matches!( - parse(br#"{"v":1,"type":"nope"}"#).unwrap_err(), - EngineStatsParseError::UnknownType(_) - )); - assert!(matches!( - parse(br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama"}"#).unwrap_err(), - EngineStatsParseError::EmptyStatsCounters - )); - assert!(matches!( - parse( - br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","tokens_processed":-1}"#, - ) - .unwrap_err(), - EngineStatsParseError::InvalidField("tokens_processed") - )); - assert!(matches!( - parse( - br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","tokens_generated":1.5}"#, + fn rejects_empty_request_updates() { + let update = request_update(proto::RequestStats { + request_id: "request-a".to_string(), + model: "model-a".to_string(), + tokens_processed: None, + tokens_generated: None, + finished: false, + }); + assert!( + translate_update( + &EngineStatsStreamConfig::default(), + update, + TokioInstant::now() ) - .unwrap_err(), - EngineStatsParseError::InvalidField("tokens_generated") - )); - - for (json, field) in [ - (br#"{"v":"1","type":"ping"}"#.as_slice(), "v"), - (br#"{"v":"\u0031","type":"ping"}"#.as_slice(), "v"), - (br#"{"v":true,"type":"ping"}"#.as_slice(), "v"), - (br#"{"v":{},"type":"ping"}"#.as_slice(), "v"), - (br#"{"v":1,"type":1}"#.as_slice(), "type"), - (br#"{"v":1,"type":-1}"#.as_slice(), "type"), - ( - br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","tokens_processed":true}"#.as_slice(), - "tokens_processed", - ), - ( - br#"{"v":1,"type":"stats","request_id":"req-1","model":"llama","finished":"true"}"#.as_slice(), - "finished", - ), - ] { - assert!(matches!( - parse(json).unwrap_err(), - EngineStatsParseError::InvalidField(actual) if actual == field - )); - } + .is_err() + ); } #[test] - fn engine_stats_parser_enforces_json_boundary_policy() { - assert!(matches!( - parse(br#"[1,2,3]"#).unwrap_err(), - EngineStatsParseError::NotObject - )); - assert!(matches!( - parse(br#"{"v":1,"type":"ping"} trailing"#).unwrap_err(), - EngineStatsParseError::Json(_) - )); - assert!(matches!( - parse(br#"{"v":1,"type":"ping","ignored":{"nested":true}}"#) - .expect("unknown fields should remain forward-compatible"), - ParsedEngineStatsEvent::Ping - )); - assert!(matches!( - parse(br#"{"v":2,"v":1,"type":"unknown","type":"ping"}"#) - .expect("recognized duplicate fields should retain last-value-wins behavior"), - ParsedEngineStatsEvent::Ping - )); - - let event = parse( - br#"{"v":1,"type":"stats","request_id":"req-\u0031","model":"ll\u0061ma","tokens_generated":1,"tokens_generated":2}"#, + fn translates_kv_snapshots_on_the_same_stream() { + let update = proto::StatsUpdate { + update: Some(proto::stats_update::Update::KvStats( + proto::KvStatsSnapshot { + snapshot_id: 1, + observed_at_unix_ms: 10, + models: Vec::new(), + }, + )), + }; + let (_, update) = translate_update( + &EngineStatsStreamConfig::default(), + update, + TokioInstant::now(), ) - .expect("escaped strings and duplicate counters should parse"); - assert!(matches!( - event, - ParsedEngineStatsEvent::Stats(update) - if update.request_id == "req-1" - && update.model_id == "llama" - && update.tokens_generated == Some(2) - )); - - assert!(matches!( - parse(br#"{"v":1,"type":"stats","request_id":null,"model":"m","finished":true}"#) - .unwrap_err(), - EngineStatsParseError::InvalidField("request_id") - )); + .unwrap(); + assert!(matches!(update, StatsAggregatorUpdate::KvCache(_))); + } + + #[derive(Clone)] + struct TestFrontendStats { + calls: Arc, + updates: Arc>, + } + + #[tonic::async_trait] + impl proto::frontend_stats_server::FrontendStats for TestFrontendStats { + type WatchStatsStream = + Pin> + Send>>; + type WatchKvPlacementsStream = + Pin> + Send>>; + + async fn watch_stats( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + self.calls.fetch_add(1, Ordering::Relaxed); + let stream = futures::stream::iter(self.updates.as_ref().clone().into_iter().map(Ok)); + Ok(tonic::Response::new(Box::pin(stream))) + } - let nested_values = "0,".repeat(16_000); - let nested_invalid = format!( - r#"{{"v":1,"type":"stats","request_id":"req-1","model":"m","tokens_generated":[{nested_values}0]}}"# - ); - assert!(matches!( - parse(nested_invalid.as_bytes()).unwrap_err(), - EngineStatsParseError::InvalidField("tokens_generated") - )); + async fn watch_kv_placements( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented("not used")) + } } - #[tokio::test] - async fn next_engine_stats_chunk_reads_before_cancellation() { - let config = EngineStatsStreamConfig::default(); - let stop = CancellationToken::new(); - let (chunk_tx, chunk_rx) = tokio::sync::mpsc::channel(1); - let reader = tokio::spawn(async move { - let mut stream = tokio_stream::wrappers::ReceiverStream::new(chunk_rx); - next_engine_stats_chunk(&config, &stop, &mut stream).await + async fn spawn_test_server(router: axum::Router) -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let endpoint = format!("http://{}", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { + axum::serve(listener, router).await.unwrap(); }); - tokio::task::yield_now().await; - chunk_tx - .send(Ok::<_, reqwest::Error>(Bytes::from_static(b"chunk"))) - .await - .expect("chunk receiver should stay alive"); - let chunk = match reader.await.expect("reader task should join") { - Ok(chunk) => chunk, - Err(_) => panic!("uncancelled token should not stop chunk reads"), - }; - - assert_eq!(chunk, Bytes::from_static(b"chunk")); + (endpoint, server) } #[tokio::test] - async fn stats_update_send_wakes_on_cancellation_when_channel_is_full() { - let (tx, _rx) = flume::bounded(1); - tx.send(StatsAggregatorUpdate::EnableOpenAiFallback) - .expect("seed update should fill channel"); - let stop = CancellationToken::new(); - let task_stop = stop.clone(); - - let task = tokio::spawn(async move { - send_stats_update(&tx, StatsAggregatorUpdate::EnableOpenAiFallback, &task_stop).await + async fn grpc_stream_delivers_both_updates_and_reconnects_after_eof() { + let calls = Arc::new(AtomicUsize::new(0)); + let updates = Arc::new(vec![ + request_update(proto::RequestStats { + request_id: "request-a".to_string(), + model: "model-a".to_string(), + tokens_processed: Some(10), + tokens_generated: None, + finished: false, + }), + proto::StatsUpdate { + update: Some(proto::stats_update::Update::KvStats( + proto::KvStatsSnapshot { + snapshot_id: 1, + observed_at_unix_ms: 10, + models: Vec::new(), + }, + )), + }, + ]); + let service = proto::frontend_stats_server::FrontendStatsServer::new(TestFrontendStats { + calls: calls.clone(), + updates, }); - tokio::task::yield_now().await; - stop.cancel(); + let router = tonic::service::Routes::new(service).into_axum_router(); + let (endpoint, server) = spawn_test_server(router).await; + let config = EngineStatsStreamConfig { + initial_reconnect_backoff: Duration::from_millis(1), + max_reconnect_backoff: Duration::from_millis(1), + ..EngineStatsStreamConfig::new(&endpoint, EngineStatsStreamMode::Required) + }; + let (tx, rx) = flume::bounded(8); + let stream = start_engine_stats_stream(config, tx).unwrap(); - let sent = tokio::time::timeout(Duration::from_secs(1), task) + let first = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) .await - .expect("stats update send should wake on cancellation") - .expect("stats update send should not panic"); - assert!(!sent); - } - - #[tokio::test] - async fn empty_engine_stats_chunks_force_reconnect_without_spinning() { - let config = EngineStatsStreamConfig::default(); - let stop = CancellationToken::new(); - let mut empty_stream = futures::stream::iter([Ok::<_, reqwest::Error>(Bytes::new())]); - - assert!(matches!( - next_engine_stats_chunk(&config, &stop, &mut empty_stream).await, - Err(StreamReadOutcome::Retry("empty_chunk")) - )); - } - - #[tokio::test] - async fn engine_stats_line_length_allows_exact_limit_only() { - let line = - br#"{"v":1,"type":"stats","request_id":"req-1","model":"model-a","tokens_generated":1}"#; - let mut config = EngineStatsStreamConfig { - max_line_bytes: line.len(), - ..Default::default() - }; - process_lines(&config, [line.as_slice(), b"\n"]) + .unwrap() + .unwrap(); + let second = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) .await - .assert_counter("exact-limit line should publish stats"); - - process_lines( - &config, - [br#"{"v":1,"type":"ping"}"#.as_slice(), b"\n", line, b"\n"], - ) - .await - .assert_counter("exact-limit line should publish after a prior line"); - - config.max_line_bytes = line.len() - 1; - process_lines(&config, [line.as_slice(), b"\n"]) + .unwrap() + .unwrap(); + let third = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) .await - .assert_no_update(false); - - let metrics = PylonMetrics::new().expect("metrics should initialize"); - config.metrics = Some(metrics.clone()); - process_lines( - &config, - [line.as_slice(), b"\n", br#"{"v":1,"type":"ping"}"#, b"\n"], - ) - .await - .assert_no_update(true); - let body = metrics.gather_text().expect("metrics should encode"); - assert!(body.contains( - r#"pylon_engine_stats_stream_invalid_events_total{reason="line_too_large"} 1"# - )); - assert!(body.contains(r#"pylon_engine_stats_stream_events_total{type="ping"} 1"#)); - assert!(!body.contains(r#"pylon_engine_stats_stream_invalid_events_total{reason="json"}"#)); - - let suffix = br#"{"v":1,"type":"stats","request_id":"suffix","model":"model-a","tokens_generated":1}"#; - let later = br#"{"v":1,"type":"stats","request_id":"later","model":"model-a","tokens_generated":1}"#; - let max_line_bytes = suffix.len(); - let oversized_prefix = vec![b'x'; max_line_bytes + 1]; - let discard_then_resume = [suffix.as_slice(), b"\n", later.as_slice(), b"\n"].concat(); - for chunks in [ - vec![oversized_prefix.as_slice(), discard_then_resume.as_slice()], - vec![ - &oversized_prefix[..max_line_bytes], - &oversized_prefix[max_line_bytes..], - &suffix[..7], - &suffix[7..], - b"\n", - later.as_slice(), - b"\n", - ], - ] { - let metrics = PylonMetrics::new().expect("metrics should initialize"); - let config = EngineStatsStreamConfig { - max_line_bytes, - metrics: Some(metrics.clone()), - ..Default::default() - }; - let processed = process_lines(&config, chunks).await; - processed.assert_state(true); - let updates = processed.updates.try_iter().collect::>(); - assert!(matches!( - updates.as_slice(), - [StatsAggregatorUpdate::RequestCounters(update)] if update.request_id == "later" - )); - let body = metrics.gather_text().expect("metrics should encode"); - assert!(body.contains( - r#"pylon_engine_stats_stream_invalid_events_total{reason="line_too_large"} 1"# - )); - assert!(body.contains(r#"pylon_engine_stats_stream_events_total{type="stats"} 1"#)); - } - } - - #[tokio::test] - async fn blank_lf_and_crlf_engine_stats_lines_are_ignored() { - let metrics = PylonMetrics::new().expect("metrics should initialize"); - let config = EngineStatsStreamConfig { - metrics: Some(metrics.clone()), - ..Default::default() - }; - process_lines( - &config, - [concat!( - "\n\r\n", - r#"{"v":1,"type":"stats","request_id":"req-1","model":"model-a","tokens_generated":1}"#, - "\r\n" - ) - .as_bytes()], - ) - .await - .assert_counter("stats line after blank CRLF should publish"); - let body = metrics.gather_text().expect("metrics should encode"); - assert!(!body.contains(r#"pylon_engine_stats_stream_invalid_events_total{reason="json"}"#)); - } - - #[test] - fn compact_engine_stats_line_buffer_keeps_partial_tail() { - let config = EngineStatsStreamConfig::default(); - let mut line_buffer = b"{\"v\":1,\"type\":\"ping\"}\n{\"v\":1".to_vec(); - assert!(!compact_line_buffer(&config, &mut line_buffer, 22)); - - assert_eq!(line_buffer, br#"{"v":1"#); - assert!(!compact_line_buffer(&config, &mut line_buffer, 6)); - assert!(line_buffer.is_empty()); - - let config = EngineStatsStreamConfig { - max_line_bytes: 4, - ..Default::default() - }; - let mut line_buffer = b"1234".to_vec(); - assert!(!compact_line_buffer(&config, &mut line_buffer, 0)); - assert_eq!(line_buffer, b"1234"); - - let mut line_buffer = b"12345".to_vec(); - assert!(compact_line_buffer(&config, &mut line_buffer, 0)); - assert!(line_buffer.is_empty()); + .unwrap() + .unwrap(); + assert!(matches!(first, StatsAggregatorUpdate::RequestCounters(_))); + assert!(matches!(second, StatsAggregatorUpdate::KvCache(_))); + assert!(matches!(third, StatsAggregatorUpdate::RequestCounters(_))); + assert!(calls.load(Ordering::Relaxed) >= 2); + + stream.shutdown().await; + server.abort(); } #[tokio::test] - async fn auto_mode_enables_openai_fallback_when_endpoint_is_unsupported_before_events() { - let (base_url, server) = serve_stats(Router::new().route( - "/pylon/v1/stats/stream", - get(|| async { StatusCode::NOT_FOUND }), - )) - .await; - + async fn auto_mode_falls_back_when_grpc_service_is_absent() { + let (endpoint, server) = spawn_test_server(axum::Router::new()).await; + let config = EngineStatsStreamConfig::new(&endpoint, EngineStatsStreamMode::Auto); let (tx, rx) = flume::bounded(1); - let handle = start_engine_stats_stream( - reconnecting_config(&base_url, EngineStatsStreamMode::Auto), - tx, - ) - .expect("auto stats stream should start"); - let update = receive_update(&rx, "auto mode should enable fallback").await; + let stream = start_engine_stats_stream(config, tx).unwrap(); + let update = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) + .await + .expect("auto mode should resolve unsupported gRPC") + .unwrap(); assert!(matches!( update, StatsAggregatorUpdate::EnableOpenAiFallback )); - handle.shutdown().await; - server.abort(); - } - - #[tokio::test] - async fn required_mode_does_not_enable_openai_fallback_for_unsupported_endpoint() { - let (base_url, server) = serve_stats(Router::new().route( - "/pylon/v1/stats/stream", - get(|| async { StatusCode::NOT_FOUND }), - )) - .await; - - let (tx, rx) = flume::bounded(1); - let handle = start_engine_stats_stream( - reconnecting_config(&base_url, EngineStatsStreamMode::Required), - tx, - ) - .expect("required stats stream should start"); - assert!( - tokio::time::timeout(Duration::from_millis(50), rx.recv_async()) - .await - .is_err(), - "required mode must not enable OpenAI fallback" - ); - - handle.shutdown().await; + stream.shutdown().await; server.abort(); } - - #[tokio::test] - async fn auto_mode_retries_unsupported_endpoint_after_valid_event_without_fallback() { - let attempts = Arc::new(AtomicUsize::new(0)); - let server_attempts = attempts.clone(); - let (base_url, server) = serve_stats( - Router::new().route( - "/pylon/v1/stats/stream", - get(move || { - let attempts = server_attempts.clone(); - async move { - if attempts.fetch_add(1, Ordering::SeqCst) == 0 { - ( - StatusCode::OK, - "{\"v\":1,\"type\":\"stats\",\"request_id\":\"req-1\",\"model\":\"model-a\",\"tokens_generated\":1}\n", - ) - } else { - (StatusCode::NOT_FOUND, "") - } - } - }), - ), - ) - .await; - - let (tx, rx) = flume::bounded(4); - let handle = start_engine_stats_stream( - reconnecting_config(&base_url, EngineStatsStreamMode::Auto), - tx, - ) - .expect("auto stats stream should start"); - let update = receive_update(&rx, "valid stats event should be sent").await; - assert!(matches!(update, StatsAggregatorUpdate::RequestCounters(_))); - assert!( - tokio::time::timeout(Duration::from_millis(50), rx.recv_async()) - .await - .is_err(), - "auto mode must not switch to fallback after any valid stream event" - ); - - handle.shutdown().await; - server.abort(); - } - - #[tokio::test] - async fn stream_invalid_events_record_metric_reasons_before_valid_update() { - let (base_url, server) = serve_stats( - Router::new().route( - "/pylon/v1/stats/stream", - get(|| async { - concat!( - "not-json\n", - "[1]\n", - "{\"type\":\"ping\"}\n", - "{\"v\":2,\"type\":\"ping\"}\n", - "{\"v\":1,\"type\":\"nope\"}\n", - "{\"v\":1,\"type\":\"stats\",\"request_id\":\"req-empty\",\"model\":\"model-a\"}\n", - "{\"v\":1,\"type\":\"stats\",\"request_id\":\"\",\"model\":\"model-a\",\"tokens_generated\":1}\n", - "{\"v\":1,\"type\":\"stats\",\"request_id\":\"req-valid\",\"model\":\"model-a\",\"tokens_generated\":1}\n", - ) - }), - ), - ) - .await; - - let metrics = PylonMetrics::new().expect("metrics should initialize"); - let (tx, rx) = flume::bounded(4); - let mut config = EngineStatsStreamConfig::new( - &base_url, - "/pylon/v1/stats/stream", - EngineStatsStreamMode::Required, - ); - config.initial_reconnect_backoff = Duration::from_secs(60); - config.max_reconnect_backoff = Duration::from_secs(60); - config.metrics = Some(metrics.clone()); - - let handle = - start_engine_stats_stream(config, tx).expect("required stats stream should start"); - let update = receive_update(&rx, "valid stats event should be sent").await; - let StatsAggregatorUpdate::RequestCounters(update) = update else { - panic!("expected valid stream line to produce request counters"); - }; - assert_eq!(update.request_id, "req-valid"); - assert_eq!(update.tokens_generated, Some(1)); - - handle.shutdown().await; - server.abort(); - - let body = metrics.gather_text().expect("metrics should encode"); - for reason in [ - "json", - "not_object", - "missing_field", - "version", - "type", - "empty_stats", - "field", - ] { - assert!( - body.contains(&format!( - r#"pylon_engine_stats_stream_invalid_events_total{{reason="{reason}"}} 1"# - )), - "missing invalid-event metric for reason {reason}; body:\n{body}" - ); - } - assert!(body.contains(r#"pylon_engine_stats_stream_events_total{type="stats"} 1"#)); - } } diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs new file mode 100644 index 000000000..9e0d734ba --- /dev/null +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs @@ -0,0 +1,113 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Translation of Dynamo frontend KV snapshots into Pylon's aggregate model state. + +use std::collections::HashSet; + +use stargate_proto::dynamo_frontend_stats as proto; + +use super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot}; + +pub(super) fn kv_snapshot_from_proto( + snapshot: proto::KvStatsSnapshot, +) -> anyhow::Result { + let mut identities = HashSet::new(); + let models = snapshot + .models + .into_iter() + .map(|model| -> anyhow::Result<_> { + anyhow::ensure!(!model.model.trim().is_empty(), "KV stats model is empty"); + anyhow::ensure!( + identities.insert(model.model.clone()), + "duplicate KV stats identity {}", + model.model + ); + for alias in &model.aliases { + anyhow::ensure!(!alias.trim().is_empty(), "KV stats alias is empty"); + anyhow::ensure!( + identities.insert(alias.clone()), + "duplicate KV stats identity {alias}" + ); + } + + let complete = snapshot.observed_at_unix_ms > 0 + && model.routing_cache.as_ref().is_some_and(|routing| { + routing.capacity_tokens > 0 + && routing.used_tokens.checked_add(routing.free_tokens) + == Some(routing.capacity_tokens) + }); + let (capacity, used, free) = model + .routing_cache + .map(|routing| { + ( + routing.capacity_tokens, + routing.used_tokens, + routing.free_tokens, + ) + }) + .unwrap_or_default(); + Ok(KvCacheStatsSnapshot { + model: model.model, + aliases: model.aliases, + kv_cache_capacity_tokens: capacity, + kv_cache_used_tokens: used, + kv_cache_free_tokens: free, + source_observed_at_unix_ms: snapshot.observed_at_unix_ms, + complete, + }) + }) + .collect::>>()?; + Ok(KvCacheStatsEnvelope { models }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn snapshot(capacity: u64, used: u64, free: u64) -> proto::KvStatsSnapshot { + proto::KvStatsSnapshot { + snapshot_id: 1, + observed_at_unix_ms: 10, + models: vec![proto::ModelKvStats { + model: "model-a".to_string(), + aliases: vec!["alias-a".to_string()], + routing_cache: Some(proto::RoutingCacheStats { + role: proto::WorkerRole::Aggregated as i32, + capacity_tokens: capacity, + used_tokens: used, + free_tokens: free, + }), + pools: Vec::new(), + }], + } + } + + #[test] + fn converts_complete_routing_cache_stats() { + let envelope = kv_snapshot_from_proto(snapshot(100, 40, 60)).unwrap(); + let model = &envelope.models[0]; + assert!(model.complete); + assert_eq!(model.kv_cache_capacity_tokens, 100); + assert_eq!(model.kv_cache_used_tokens, 40); + assert_eq!(model.kv_cache_free_tokens, 60); + } + + #[test] + fn marks_inconsistent_totals_incomplete() { + let envelope = kv_snapshot_from_proto(snapshot(100, 40, 50)).unwrap(); + assert!(!envelope.models[0].complete); + } + + #[test] + fn rejects_duplicate_model_identity() { + let mut value = snapshot(100, 40, 60); + value.models.push(proto::ModelKvStats { + model: "alias-a".to_string(), + aliases: Vec::new(), + routing_cache: None, + pools: Vec::new(), + }); + assert!(kv_snapshot_from_proto(value).is_err()); + } +} diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs deleted file mode 100644 index d0b759111..000000000 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats_stream.rs +++ /dev/null @@ -1,377 +0,0 @@ -// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -// SPDX-License-Identifier: Apache-2.0 - -use std::collections::HashSet; -use std::time::Duration; - -use bytes::Bytes; -use futures::{Stream, StreamExt}; -use serde::Deserialize; -use tokio_util::sync::CancellationToken; - -use super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot}; - -const MAX_LINE_BYTES: usize = 1024 * 1024; - -pub(super) struct KvStatsStreamConfig { - pub(super) url: String, - pub(super) reconnect_interval: Duration, - pub(super) connect_timeout: Duration, - pub(super) idle_timeout: Duration, -} - -#[derive(Deserialize)] -struct RawSnapshot { - v: u8, - #[serde(rename = "type")] - event_type: String, - observed_at_unix_ms: u64, - models: Vec, -} - -#[derive(Deserialize)] -struct RawModelStats { - model: String, - #[serde(default)] - aliases: Vec, - routing_cache: Option, -} - -#[derive(Deserialize)] -struct RawRoutingCacheStats { - capacity_tokens: Option, - used_tokens: Option, - free_tokens: Option, -} - -pub(super) fn parse_kv_stats_snapshot(line: &[u8]) -> anyhow::Result { - let snapshot: RawSnapshot = serde_json::from_slice(line)?; - anyhow::ensure!( - snapshot.v == 1, - "unsupported KV stats version {}", - snapshot.v - ); - anyhow::ensure!( - snapshot.event_type == "kv_stats_snapshot", - "unsupported KV stats event type {}", - snapshot.event_type - ); - let mut identities = HashSet::new(); - let models = snapshot - .models - .into_iter() - .map(|model| -> anyhow::Result<_> { - anyhow::ensure!(!model.model.trim().is_empty(), "KV stats model is empty"); - anyhow::ensure!( - identities.insert(model.model.clone()), - "duplicate KV stats identity {}", - model.model - ); - for alias in &model.aliases { - anyhow::ensure!(!alias.trim().is_empty(), "KV stats alias is empty"); - anyhow::ensure!( - identities.insert(alias.clone()), - "duplicate KV stats identity {alias}" - ); - } - let routing = model.routing_cache; - let complete = snapshot.observed_at_unix_ms > 0 - && routing.as_ref().is_some_and(|routing| { - let Some((capacity, used, free)) = routing - .capacity_tokens - .zip(routing.used_tokens) - .zip(routing.free_tokens) - .map(|((capacity, used), free)| (capacity, used, free)) - else { - return false; - }; - capacity > 0 && used.checked_add(free) == Some(capacity) - }); - let (capacity, used, free) = routing - .map(|routing| { - ( - routing.capacity_tokens.unwrap_or_default(), - routing.used_tokens.unwrap_or_default(), - routing.free_tokens.unwrap_or_default(), - ) - }) - .unwrap_or_default(); - Ok(KvCacheStatsSnapshot { - model: model.model, - aliases: model.aliases, - kv_cache_capacity_tokens: capacity, - kv_cache_used_tokens: used, - kv_cache_free_tokens: free, - source_observed_at_unix_ms: snapshot.observed_at_unix_ms, - complete, - }) - }) - .collect::>>()?; - Ok(KvCacheStatsEnvelope { models }) -} - -pub(super) async fn run_kv_stats_stream( - config: KvStatsStreamConfig, - updates: flume::Sender, - stop: CancellationToken, -) { - let client = reqwest::Client::new(); - while !stop.is_cancelled() { - if let Err(error) = read_stream_once(&config, &client, &updates, &stop).await { - tracing::warn!(url = config.url, %error, "KV stats stream disconnected"); - } - if stop - .run_until_cancelled(tokio::time::sleep(config.reconnect_interval)) - .await - .is_none() - { - break; - } - } -} - -async fn read_stream_once( - config: &KvStatsStreamConfig, - client: &reqwest::Client, - updates: &flume::Sender, - stop: &CancellationToken, -) -> anyhow::Result<()> { - let response = tokio::select! { - _ = stop.cancelled() => return Ok(()), - response = tokio::time::timeout( - config.connect_timeout, - client - .get(&config.url) - .header(reqwest::header::ACCEPT, "application/x-ndjson") - .send(), - ) => response??, - }; - anyhow::ensure!( - response.status().is_success(), - "KV stats endpoint returned {}", - response.status() - ); - drain_response(response.bytes_stream(), updates, stop, config.idle_timeout).await -} - -async fn drain_response( - mut stream: S, - updates: &flume::Sender, - stop: &CancellationToken, - idle_timeout: Duration, -) -> anyhow::Result<()> -where - S: Stream> + Unpin, -{ - let mut buffer = Vec::with_capacity(4096); - let mut discarding_oversized_line = false; - loop { - let chunk = tokio::select! { - _ = stop.cancelled() => return Ok(()), - chunk = tokio::time::timeout(idle_timeout, stream.next()) => { - chunk.map_err(|_| anyhow::anyhow!("KV stats stream became idle"))? - }, - }; - let Some(chunk) = chunk else { - anyhow::bail!("KV stats stream ended"); - }; - let chunk = chunk?; - let mut remaining = chunk.as_ref(); - while let Some(newline) = remaining.iter().position(|byte| *byte == b'\n') { - let segment = &remaining[..newline]; - remaining = &remaining[newline + 1..]; - if discarding_oversized_line { - discarding_oversized_line = false; - continue; - } - if buffer.len().saturating_add(segment.len()) > MAX_LINE_BYTES { - tracing::warn!("dropping oversized KV stats line"); - buffer.clear(); - continue; - } - buffer.extend_from_slice(segment); - if buffer.iter().all(u8::is_ascii_whitespace) { - buffer.clear(); - continue; - } - match parse_kv_stats_snapshot(&buffer) { - Ok(snapshot) => { - match stop.run_until_cancelled(updates.send_async(snapshot)).await { - None | Some(Err(_)) => return Ok(()), - Some(Ok(())) => {} - } - } - Err(error) => tracing::warn!(%error, "dropping invalid KV stats snapshot"), - } - buffer.clear(); - } - if discarding_oversized_line { - continue; - } - if buffer.len().saturating_add(remaining.len()) > MAX_LINE_BYTES { - tracing::warn!("dropping oversized KV stats line"); - buffer.clear(); - discarding_oversized_line = true; - } else { - buffer.extend_from_slice(remaining); - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn routing_cache_presence_denotes_a_complete_snapshot() { - let snapshot = parse_kv_stats_snapshot( - br#"{"v":1,"type":"kv_stats_snapshot","snapshot_id":9,"observed_at_unix_ms":42,"models":[{"model":"m","aliases":["alias"],"routing_cache":{"role":"decode","capacity_tokens":100,"used_tokens":40,"free_tokens":60},"pools":[]}]}"#, - ) - .unwrap(); - assert_eq!(snapshot.models[0].source_observed_at_unix_ms, 42); - assert_eq!(snapshot.models.len(), 1); - assert_eq!(snapshot.models[0].aliases, ["alias"]); - assert!(snapshot.models[0].complete); - } - - #[test] - fn inconsistent_routing_cache_snapshot_is_not_usable() { - let snapshot = parse_kv_stats_snapshot( - br#"{"v":1,"type":"kv_stats_snapshot","observed_at_unix_ms":42,"models":[{"model":"m","aliases":[],"routing_cache":{"capacity_tokens":100,"used_tokens":80,"free_tokens":30},"pools":[]}]}"#, - ) - .unwrap(); - assert!(!snapshot.models[0].complete); - } - - #[test] - fn duplicate_alias_ownership_rejects_the_whole_snapshot() { - let result = parse_kv_stats_snapshot( - br#"{"v":1,"type":"kv_stats_snapshot","observed_at_unix_ms":42,"models":[{"model":"a","aliases":["shared"],"routing_cache":null},{"model":"b","aliases":["shared"],"routing_cache":null}]}"#, - ); - assert!(result.is_err()); - } - - #[test] - fn line_limit_accepts_a_representative_thousand_model_snapshot() { - let models = (0..1_000) - .map(|index| { - serde_json::json!({ - "model": format!("model-{index:04}"), - "aliases": [format!("deployment-model-{index:04}")], - "routing_cache": { - "role": "decode", - "capacity_tokens": 65_536_000, - "used_tokens": 6_553_600, - "free_tokens": 58_982_400 - }, - "pools": [{ - "namespace": "dynamo", - "component": "backend", - "endpoint": "generate", - "role": "decode", - "storage_tier": "device", - "block_size_tokens": 64, - "expected_ranks": 8, - "observed_ranks": 8, - "capacity_blocks": 1_024_000, - "used_blocks": 102_400, - "free_blocks": 921_600, - "active_decode_blocks": 81_920, - "complete": true - }] - }) - }) - .collect::>(); - let line = serde_json::to_vec(&serde_json::json!({ - "v": 1, - "type": "kv_stats_snapshot", - "snapshot_id": 1, - "observed_at_unix_ms": 1, - "models": models - })) - .unwrap(); - - assert!( - line.len() <= MAX_LINE_BYTES, - "representative snapshot is {} bytes", - line.len() - ); - } - - #[tokio::test] - async fn fragmented_ndjson_is_reassembled_before_publication() { - let chunks = futures::stream::iter([ - Ok::<_, reqwest::Error>(Bytes::from_static( - b"{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":9,", - )), - Ok(Bytes::from_static( - b"\"models\":[{\"model\":\"m\",\"routing_cache\":{\"capacity_tokens\":10,\"used_tokens\":4,\"free_tokens\":6}}]}\n", - )), - ]); - let (tx, rx) = flume::bounded(1); - let result = drain_response( - chunks, - &tx, - &CancellationToken::new(), - Duration::from_secs(1), - ) - .await; - assert!( - result.is_err(), - "finite response should end after publishing" - ); - let snapshot = rx.try_recv().expect("snapshot should be published"); - assert_eq!(snapshot.models[0].source_observed_at_unix_ms, 9); - assert!(snapshot.models[0].complete); - } - - #[tokio::test] - async fn oversized_line_is_dropped_without_buffering_or_losing_the_next_snapshot() { - let valid = - b"{\"v\":1,\"type\":\"kv_stats_snapshot\",\"observed_at_unix_ms\":9,\"models\":[]}\n"; - let mut first = vec![b'x'; MAX_LINE_BYTES + 1]; - first.extend_from_slice(b"\n"); - let chunks = futures::stream::iter([ - Ok::<_, reqwest::Error>(Bytes::from(first)), - Ok(Bytes::from_static(valid)), - ]); - let (tx, rx) = flume::bounded(1); - - let result = drain_response( - chunks, - &tx, - &CancellationToken::new(), - Duration::from_secs(1), - ) - .await; - - assert!( - result.is_err(), - "finite response should end after publishing" - ); - assert_eq!( - rx.try_recv() - .expect("valid line should still publish") - .models - .len(), - 0 - ); - } - - #[tokio::test] - async fn idle_stream_is_reconnected() { - let stream = futures::stream::pending::>(); - let (tx, _rx) = flume::bounded(1); - - let error = drain_response( - stream, - &tx, - &CancellationToken::new(), - Duration::from_millis(10), - ) - .await - .unwrap_err(); - - assert!(error.to_string().contains("became idle")); - } -} diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs index 020e4471b..dc3d5c750 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/mod.rs @@ -36,7 +36,7 @@ macro_rules! owned_task_handle { mod aggregator; mod collector; mod engine_stats_stream; -mod kv_stats_stream; +mod kv_stats; mod metrics; mod projection; pub(crate) mod token_metrics; @@ -49,7 +49,7 @@ pub use collector::{ }; pub use engine_stats_stream::{ EngineStatsStreamConfig, EngineStatsStreamHandle, EngineStatsStreamMode, - parse_engine_stats_line_for_benchmark, start_engine_stats_stream, + start_engine_stats_stream, }; pub(crate) use metrics::CalibrationOutcome; pub use metrics::{MetricsServerHandle, PylonMetrics, start_metrics_server}; diff --git a/src/libraries/rust/stargate/crates/pylon/src/main.rs b/src/libraries/rust/stargate/crates/pylon/src/main.rs index ff64f84fa..eec83d7aa 100644 --- a/src/libraries/rust/stargate/crates/pylon/src/main.rs +++ b/src/libraries/rust/stargate/crates/pylon/src/main.rs @@ -110,15 +110,9 @@ struct Args { /// Timeout for calibration requests in milliseconds #[arg(long, default_value_t = 30000, value_name = "MS")] bringup_calibration_timeout_ms: u64, - /// Upstream HTTP path for the canonical KV-cache stats stream - #[arg(long, default_value = "/v1/kv-cache/stats/stream", value_name = "PATH")] - kv_cache_stats_path: Option, /// Engine stats stream source selection mode #[arg(long, default_value_t = EngineStatsStreamMode::Auto, value_name = "MODE")] engine_stats_stream: EngineStatsStreamMode, - /// Upstream HTTP path for the engine stats stream - #[arg(long, default_value = "/pylon/v1/stats/stream", value_name = "PATH")] - engine_stats_stream_path: String, /// Keep --initial-input-tps fixed for deterministic benchmark/test experiments #[arg(long, default_value_t = false, hide = true)] benchmark_pin_input_tps: bool, @@ -278,7 +272,7 @@ mod tests { use reqwest::header::HeaderName; use super::startup::{ - effective_cluster_id, normalize_base_url, pylon_queue_mismatch_retry_config_from_args, + effective_cluster_id, pylon_queue_mismatch_retry_config_from_args, pylon_retry_config_from_args, request_quality_monitor_config_from_args, stats_collector_config_from_args, }; @@ -580,17 +574,11 @@ mod tests { } #[test] - fn engine_stats_stream_defaults_to_auto_mode_and_v1_path() { + fn engine_stats_stream_defaults_to_auto_mode() { let args = parse_args(""); - let upstream = normalize_base_url(&args.upstream_http_base_url); - let metrics_config = stats_collector_config_from_args(&args, &upstream); + let metrics_config = stats_collector_config_from_args(&args); assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Auto); - assert_eq!(args.engine_stats_stream_path, "/pylon/v1/stats/stream"); - assert_eq!( - metrics_config.kv_cache_stats_url.as_deref(), - Some("http://127.0.0.1:8090/v1/kv-cache/stats/stream") - ); assert!( !metrics_config.openai_fallback_stats_enabled, "auto mode should wait for a permanent unsupported stream response before fallback stats" @@ -600,34 +588,16 @@ mod tests { #[test] fn engine_stats_stream_can_be_disabled() { let args = parse_args("--engine-stats-stream off"); - let upstream = normalize_base_url(&args.upstream_http_base_url); - let metrics_config = stats_collector_config_from_args(&args, &upstream); + let metrics_config = stats_collector_config_from_args(&args); assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Off); - assert_eq!( - metrics_config.kv_cache_stats_url.as_deref(), - Some("http://127.0.0.1:8090/v1/kv-cache/stats/stream") - ); assert!(metrics_config.openai_fallback_stats_enabled); } - #[test] - fn kv_cache_stats_path_overrides_the_canonical_stream() { - let args = parse_args("--kv-cache-stats-path /kv-cache/stats"); - let upstream = normalize_base_url(&args.upstream_http_base_url); - let metrics_config = stats_collector_config_from_args(&args, &upstream); - - assert_eq!( - metrics_config.kv_cache_stats_url, - Some("http://127.0.0.1:8090/kv-cache/stats".to_string()) - ); - } - #[test] fn required_engine_stats_stream_disables_openai_fallback_stats() { let args = parse_args("--engine-stats-stream required"); - let upstream = normalize_base_url(&args.upstream_http_base_url); - let metrics_config = stats_collector_config_from_args(&args, &upstream); + let metrics_config = stats_collector_config_from_args(&args); assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Required); assert!(!metrics_config.openai_fallback_stats_enabled); diff --git a/src/libraries/rust/stargate/crates/pylon/src/startup.rs b/src/libraries/rust/stargate/crates/pylon/src/startup.rs index 6d99fe1fb..c8a1cd3aa 100644 --- a/src/libraries/rust/stargate/crates/pylon/src/startup.rs +++ b/src/libraries/rust/stargate/crates/pylon/src/startup.rs @@ -342,7 +342,7 @@ async fn start_pylon_runtime(args: &Args, plan: &PylonStartupPlan) -> Result, )> { let (stats_update_tx, stats_update_rx) = stats_aggregator_update_channel(stats_config); - let mut config = EngineStatsStreamConfig::new( - &plan.upstream, - &args.engine_stats_stream_path, - args.engine_stats_stream, - ); + let mut config = EngineStatsStreamConfig::new(&plan.upstream, args.engine_stats_stream); config.metrics = Some(metrics); config.runtime_state = Some(runtime_state); let mode = config.mode; @@ -595,20 +591,9 @@ fn tunnel_forwarding_config_from_plan( } } -pub(crate) fn stats_collector_config_from_args( - args: &Args, - upstream: &str, -) -> StatsCollectorConfig { +pub(crate) fn stats_collector_config_from_args(args: &Args) -> StatsCollectorConfig { StatsCollectorConfig { openai_fallback_stats_enabled: args.engine_stats_stream == EngineStatsStreamMode::Off, - // Dynamo exposes this canonical stream independently from request stats. - kv_cache_stats_url: args.kv_cache_stats_path.as_deref().map(|path| { - format!( - "{}/{}", - upstream.trim_end_matches('/'), - path.trim_start_matches('/') - ) - }), ..Default::default() } } @@ -1583,19 +1568,10 @@ mod tests { } #[test] - fn stats_config_uses_normalized_upstream() { - let (args, plan) = startup(&[ - "--engine-stats-stream", - "required", - "--kv-cache-stats-path", - "kv/live", - ]); - let stats = stats_collector_config_from_args(&args, &plan.upstream); + fn required_stats_config_disables_fallback() { + let (args, _) = startup(&["--engine-stats-stream", "required"]); + let stats = stats_collector_config_from_args(&args); - assert_eq!( - stats.kv_cache_stats_url.as_deref(), - Some("http://127.0.0.1:8090/kv/live") - ); assert!(!stats.openai_fallback_stats_enabled); } @@ -1603,7 +1579,7 @@ mod tests { async fn configured_input_tps_seeds_queue_estimates_before_engine_stats() { let (args, plan) = startup(&[]); let metrics = PylonMetrics::new().expect("metrics should initialize"); - let mut config = stats_collector_config_from_args(&args, &plan.upstream); + let mut config = stats_collector_config_from_args(&args); config.openai_fallback_stats_enabled = true; let (runtime_state, request_observation_rx) = PylonRuntimeState::observed( InferenceServerStatus::Unknown, diff --git a/src/libraries/rust/stargate/crates/stargate-bench/src/k8s/render.rs b/src/libraries/rust/stargate/crates/stargate-bench/src/k8s/render.rs index 4630a044e..412f95553 100644 --- a/src/libraries/rust/stargate/crates/stargate-bench/src/k8s/render.rs +++ b/src/libraries/rust/stargate/crates/stargate-bench/src/k8s/render.rs @@ -126,7 +126,7 @@ pub(super) fn render_manifest(render: RenderManifestConfig<'_>) -> RenderedManif )); } backends.push_str(&format!( - "apiVersion: apps/v1\nkind: Deployment\nmetadata:\n name: {inference_server_id}-pylon\n namespace: {backends_ns}\nspec:\n replicas: 1\n selector:\n matchLabels:\n app: {inference_server_id}-pylon\n template:\n metadata:\n labels:\n app: {inference_server_id}-pylon\n benchmark.stargate/profile: {profile_name}\n spec:\n containers:\n - name: pylon\n image: {pylon_image}\n imagePullPolicy: IfNotPresent\n args:\n - --upstream-http-base-url=http://{upstream_backend_name}-http.{backends_ns}.svc.cluster.local:8090\n - --model-name={model}\n - --stargate-address=stargate.{stargate_ns}.svc.cluster.local:50071\n - --inference-server-id={inference_server_id}\n{cluster_id_arg} - --backend-connectivity=reverse\n - --quic-insecure\n - --tunnel-protocol={tunnel_protocol}\n - --kv-cache-stats-path=/kv-cache/stats\n - --min-update-interval-ms=100\n - --disable-bringup\n - --active-canary-interval-ms=0\n - --initial-input-tps={last_mean_input_tps}\n - --benchmark-pin-input-tps\n", + "apiVersion: apps/v1\nkind: Deployment\nmetadata:\n name: {inference_server_id}-pylon\n namespace: {backends_ns}\nspec:\n replicas: 1\n selector:\n matchLabels:\n app: {inference_server_id}-pylon\n template:\n metadata:\n labels:\n app: {inference_server_id}-pylon\n benchmark.stargate/profile: {profile_name}\n spec:\n containers:\n - name: pylon\n image: {pylon_image}\n imagePullPolicy: IfNotPresent\n args:\n - --upstream-http-base-url=http://{upstream_backend_name}-http.{backends_ns}.svc.cluster.local:8090\n - --model-name={model}\n - --stargate-address=stargate.{stargate_ns}.svc.cluster.local:50071\n - --inference-server-id={inference_server_id}\n{cluster_id_arg} - --backend-connectivity=reverse\n - --quic-insecure\n - --tunnel-protocol={tunnel_protocol}\n - --min-update-interval-ms=100\n - --disable-bringup\n - --active-canary-interval-ms=0\n - --initial-input-tps={last_mean_input_tps}\n - --benchmark-pin-input-tps\n", upstream_backend_name = pylon.upstream_backend_name, inference_server_id = pylon.inference_server_id, profile_name = pylon.profile_slug, diff --git a/src/libraries/rust/stargate/crates/stargate-bench/src/orchestrator.rs b/src/libraries/rust/stargate/crates/stargate-bench/src/orchestrator.rs index 185127b72..2d412c2c4 100644 --- a/src/libraries/rust/stargate/crates/stargate-bench/src/orchestrator.rs +++ b/src/libraries/rust/stargate/crates/stargate-bench/src/orchestrator.rs @@ -291,7 +291,6 @@ fn build_compose_spec( "--backend-connectivity=reverse", "--quic-insecure", "--tunnel-protocol" => config.tunnel_protocol.to_string(), - "--kv-cache-stats-path" => "/kv-cache/stats", "--min-update-interval-ms" => "100", "--disable-bringup", "--active-canary-interval-ms=0", diff --git a/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs b/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs index f0469446a..2c3547bb0 100644 --- a/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs +++ b/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs @@ -16,6 +16,7 @@ use std::collections::HashSet; use std::io::Write; use std::net::SocketAddr; +use std::pin::Pin; use std::time::Duration; use crate::common::sse::{assert_sse_done, chat_completion_contents, parse_sse_events}; @@ -25,11 +26,12 @@ use crate::common::{ reverse_registration_config, start_dummy_backend, start_dummy_inst, wait_for_inference_server_ids, wait_for_routing, wait_for_unroutable, with_proxy_headers, }; -use axum::body::{Body, Bytes}; +use axum::body::Body; use axum::extract::State; use axum::http::{HeaderMap, Response}; use axum::routing::{get, post}; use axum::{Json, Router}; +use futures::Stream; use pylon_lib::{ EngineStatsStreamConfig, EngineStatsStreamMode, InferenceServerRegistrationClient, InferenceServerRegistrationConfig, PylonRuntimeState, QuicHttpTunnelConfig, @@ -39,7 +41,7 @@ use pylon_lib::{ }; use stargate::routing::RoutingTargetKey; use stargate::test_support::StargateState; -use stargate_proto::pb::InferenceServerStatus; +use stargate_proto::{dynamo_frontend_stats as stats_proto, pb::InferenceServerStatus}; use tokio::net::TcpListener; use tokio::sync::{broadcast, watch}; @@ -128,7 +130,6 @@ async fn end_to_end_engine_stats_stream_reports_model_stats() { runtime_state: Some(runtime_state.clone()), ..EngineStatsStreamConfig::new( &format!("http://{inst_addr}"), - "/pylon/v1/stats/stream", EngineStatsStreamMode::Required, ) }, @@ -495,7 +496,7 @@ async fn reverse_tunnel_handshake_rejects_non_reverse_instance_id() { #[derive(Clone)] struct EngineStatsState { model: String, - stats_tx: broadcast::Sender, + stats_tx: broadcast::Sender, connected_tx: watch::Sender, } @@ -512,15 +513,22 @@ async fn start_engine_stats_inst( let addr = listener.local_addr().unwrap(); let (stats_tx, _) = broadcast::channel(16); let (connected_tx, connected_rx) = watch::channel(false); + let state = EngineStatsState { + model: model.to_string(), + stats_tx, + connected_tx, + }; + let grpc = tonic::service::Routes::new( + stats_proto::frontend_stats_server::FrontendStatsServer::new(EngineStatsGrpc { + state: state.clone(), + }), + ) + .into_axum_router(); let app = Router::new() .route("/v1/chat/completions", post(engine_stats_chat)) - .route("/pylon/v1/stats/stream", get(engine_stats_stream)) .route("/health", get(|| async { "ok" })) - .with_state(EngineStatsState { - model: model.to_string(), - stats_tx, - connected_tx, - }); + .with_state(state) + .merge(grpc); tokio::spawn(async move { axum::serve(listener, app).await.unwrap(); }); @@ -535,25 +543,6 @@ async fn start_engine_stats_inst( (addr, format!("quic://{tunnel_addr}"), tunnel, connected_rx) } -async fn engine_stats_stream(State(state): State) -> Response { - let _ = state.connected_tx.send(true); - let mut events = state.stats_tx.subscribe(); - let stream = async_stream::stream! { - loop { - match events.recv().await { - Ok(event) => yield Ok::(Bytes::from(event)), - Err(broadcast::error::RecvError::Lagged(_)) => continue, - Err(broadcast::error::RecvError::Closed) => break, - } - } - }; - - Response::builder() - .header("content-type", "application/x-ndjson") - .body(Body::from_stream(stream)) - .unwrap() -} - async fn engine_stats_chat( headers: HeaderMap, State(state): State, @@ -571,30 +560,8 @@ async fn engine_stats_chat( .and_then(|value| value.to_str().ok()) .expect("test proxy should send x-request-id"); let model = state.model.clone(); - send_engine_stats_event( - &state.stats_tx, - serde_json::json!({ - "v": 1, - "type": "stats", - "request_id": request_id, - "model": model, - "tokens_processed": 1, - "tokens_generated": 0, - "finished": false, - }), - ); - send_engine_stats_event( - &state.stats_tx, - serde_json::json!({ - "v": 1, - "type": "stats", - "request_id": request_id, - "model": model, - "tokens_processed": 1, - "tokens_generated": 2, - "finished": true, - }), - ); + send_engine_stats_event(&state.stats_tx, request_id, &model, Some(1), Some(0), false); + send_engine_stats_event(&state.stats_tx, request_id, &model, Some(1), Some(2), true); let data_chunk = format!( r#"{{"object":"chat.completion.chunk","model":"{model}","choices":[{{"delta":{{"content":"Hello from engine stats"}}}}]}}"# @@ -611,8 +578,72 @@ data: [DONE]\n\n" .unwrap() } -fn send_engine_stats_event(tx: &broadcast::Sender, event: serde_json::Value) { - let _ = tx.send(format!("{event}\n")); +fn send_engine_stats_event( + tx: &broadcast::Sender, + request_id: &str, + model: &str, + tokens_processed: Option, + tokens_generated: Option, + finished: bool, +) { + let _ = tx.send(stats_proto::StatsUpdate { + update: Some(stats_proto::stats_update::Update::RequestStats( + stats_proto::RequestStats { + request_id: request_id.to_string(), + model: model.to_string(), + tokens_processed, + tokens_generated, + finished, + }, + )), + }); +} + +#[derive(Clone)] +struct EngineStatsGrpc { + state: EngineStatsState, +} + +#[tonic::async_trait] +impl stats_proto::frontend_stats_server::FrontendStats for EngineStatsGrpc { + type WatchStatsStream = + Pin> + Send>>; + type WatchKvPlacementsStream = + Pin> + Send>>; + + async fn watch_stats( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + let _ = self.state.connected_tx.send(true); + let mut events = self.state.stats_tx.subscribe(); + let stream = async_stream::stream! { + yield Ok(stats_proto::StatsUpdate { + update: Some(stats_proto::stats_update::Update::KvStats( + stats_proto::KvStatsSnapshot { + snapshot_id: 1, + observed_at_unix_ms: 1, + models: Vec::new(), + }, + )), + }); + loop { + match events.recv().await { + Ok(event) => yield Ok(event), + Err(broadcast::error::RecvError::Lagged(_)) => break, + Err(broadcast::error::RecvError::Closed) => break, + } + } + }; + Ok(tonic::Response::new(Box::pin(stream))) + } + + async fn watch_kv_placements( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented("placements are not used here")) + } } async fn wait_for_engine_stats_stream_connection( diff --git a/src/libraries/rust/stargate/docs/diagrams/system-dfd.puml b/src/libraries/rust/stargate/docs/diagrams/system-dfd.puml index cafcef4ff..842148d94 100644 --- a/src/libraries/rust/stargate/docs/diagrams/system-dfd.puml +++ b/src/libraries/rust/stargate/docs/diagrams/system-dfd.puml @@ -59,7 +59,7 @@ rectangle "pylon" #e8f4fd { } rectangle "Dynamo Instance" #fdf5e6 { - rectangle "Upstream HTTP\n/v1/chat/completions\n/v1/responses\n/v1/embeddings\n/pylon/v1/stats/stream\n/health" as Upstream + rectangle "Dynamo Frontend\nOpenAI HTTP + /health\nFrontendStats gRPC\n/v1/stats/stream (debug)" as Upstream } cloud "Metrics / OTEL\nCollectors" as Obs From b6571f9460f70f510d1dabd3b81b64038181457f Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Wed, 26 Aug 2026 20:54:55 +0000 Subject: [PATCH 7/9] feat(pylon): consume KVDCRelay stats streams Signed-off-by: Barry Greengus --- .../stargate/crates/mock-dynamo/src/main.rs | 1 - .../stargate/crates/mock-dynamo/src/openai.rs | 13 +- .../crates/mock-dynamo/src/stats_stream.rs | 401 +++++++---- .../stargate/crates/mock-dynamo/src/tests.rs | 128 +--- .../proto/proto/dynamo_frontend_stats.proto | 219 ------ .../proto/proto/dynamo_kv_dc_relay.proto | 201 ++++++ .../stargate/crates/proto/src/build_plan.rs | 6 +- .../rust/stargate/crates/proto/src/lib.rs | 6 +- .../stargate/crates/pylon-lib/src/bringup.rs | 7 - .../crates/pylon-lib/src/bringup/upstream.rs | 4 +- .../pylon-lib/src/generated_request_id.rs | 1 + .../crates/pylon-lib/src/model_lifecycle.rs | 37 +- .../crates/pylon-lib/src/queue_admission.rs | 18 +- .../pylon-lib/src/quic_http_tunnel/backend.rs | 35 +- .../pylon-lib/src/quic_http_tunnel/core.rs | 37 +- .../pylon-lib/src/quic_http_tunnel/tests.rs | 88 +-- .../crates/pylon-lib/src/request_observer.rs | 70 +- .../pylon-lib/src/request_observer/tunnel.rs | 4 + .../crates/pylon-lib/src/runtime_state.rs | 159 ++++- .../crates/pylon-lib/src/stats/aggregator.rs | 161 +++-- .../crates/pylon-lib/src/stats/collector.rs | 133 ++-- .../src/stats/engine_stats_stream.rs | 619 ++++++++--------- .../crates/pylon-lib/src/stats/kv_stats.rs | 628 +++++++++++++++--- .../crates/pylon-lib/src/stats/projection.rs | 5 - .../rust/stargate/crates/pylon/src/main.rs | 17 +- .../rust/stargate/crates/pylon/src/startup.rs | 106 +-- .../stargate/tests/suite/integration.rs | 210 +++--- 27 files changed, 1953 insertions(+), 1361 deletions(-) delete mode 100644 src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto create mode 100644 src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs index c344d6c84..930428a2b 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/main.rs @@ -128,7 +128,6 @@ async fn main() -> Result<()> { .route("/v1/models", get(openai::list_models)) .route("/v1/responses", post(openai::responses)) .route("/v1/embeddings", post(openai::embeddings)) - .route("/v1/stats/stream", get(stats_stream::stats_stream)) .route("/kv-cache/stats", get(openai::kv_cache_stats)) .route( "/test-control/models/{model}", diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs index 61d637cf5..6f5e6ceb0 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/openai.rs @@ -203,7 +203,7 @@ pub(crate) async fn chat_completions( let request_id = optional_header(&headers, "x-request-id").unwrap_or_else(|| id.clone()); let cache_affinity_key = optional_header(&headers, "x-cache-affinity-key"); if stream { - state.emit_counters(&request_id, &model, 0, 0, false); + state.emit_counters(&request_id, &model, input_tokens, 0, false); } let kv_cache_access = state .process_input_with_cache(cache_affinity_key.as_deref(), input_tokens) @@ -297,9 +297,9 @@ pub(crate) async fn responses( let output_tokens = response_output_tokens(&headers, &req, state.num_tokens); let id = format!("resp-mock-{}", rand_id()); info!(id = %id, model = %model, "received responses request"); - let request_id = optional_header(&headers, "x-request-id").unwrap_or_default(); + let request_id = optional_header(&headers, "x-request-id").unwrap_or_else(|| id.clone()); let cache_affinity_key = optional_header(&headers, "x-cache-affinity-key"); - state.emit_counters(&request_id, &model, 0, 0, false); + state.emit_counters(&request_id, &model, input_tokens, 0, false); let kv_cache_access = state .process_input_with_cache(cache_affinity_key.as_deref(), input_tokens) .await; @@ -438,8 +438,11 @@ impl AppState { let _ = self.stats_events.send(StatsStreamEvent { request_id: request_id.to_string(), model: model.to_string(), - tokens_processed: Some(input_tokens as u64), - tokens_generated: output_tokens.into().map(|tokens| tokens as u64), + input_tokens: u64::try_from(input_tokens).unwrap_or(u64::MAX), + output_tokens: output_tokens + .into() + .and_then(|tokens| u64::try_from(tokens).ok()) + .unwrap_or_default(), finished, }); } diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs index 73b5cfa40..9f77ede93 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs @@ -1,184 +1,309 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -use std::convert::Infallible; +use std::collections::{BTreeMap, BTreeSet, HashMap}; use std::pin::Pin; use std::sync::atomic::Ordering; use std::time::Duration; -use axum::body::{Body, Bytes}; -use axum::extract::State; -use axum::http::{HeaderValue, StatusCode, header}; -use axum::response::{IntoResponse, Response}; use futures::Stream; -use stargate_proto::dynamo_frontend_stats as proto; +use stargate_proto::dynamo_kv_dc_relay as proto; use tokio::sync::broadcast; -use tonic::{Request, Response as GrpcResponse, Status}; +use tonic::{Request, Response, Status}; use crate::AppState; +const SNAPSHOT_INTERVAL: Duration = Duration::from_secs(1); + #[derive(Debug, Clone)] pub(crate) struct StatsStreamEvent { pub(crate) request_id: String, pub(crate) model: String, - pub(crate) tokens_processed: Option, - pub(crate) tokens_generated: Option, + pub(crate) input_tokens: u64, + pub(crate) output_tokens: u64, pub(crate) finished: bool, } pub(crate) fn grpc_router(state: AppState) -> axum::Router { - let service = - proto::frontend_stats_server::FrontendStatsServer::new(MockFrontendStats { state }); + let service = proto::kv_dc_relay_server::KvDcRelayServer::new(MockKvDcRelay { state }); tonic::service::Routes::new(service).into_axum_router() } -pub(crate) async fn stats_stream(State(state): State) -> Response { - if !state.stats_stream_enabled.load(Ordering::Relaxed) { - return StatusCode::SERVICE_UNAVAILABLE.into_response(); - } - let source = mixed_stats_stream(state); - let stream = async_stream::stream! { - futures::pin_mut!(source); - while let Some(update) = futures::StreamExt::next(&mut source).await { - yield Ok::(ndjson_event(update)); - } - }; - let mut response = Response::new(Body::from_stream(stream)); - response.headers_mut().insert( - header::CONTENT_TYPE, - HeaderValue::from_static("application/x-ndjson"), - ); - response +#[derive(Clone)] +struct MockKvDcRelay { + state: AppState, } -fn mixed_stats_stream(state: AppState) -> impl Stream + Send + 'static { - let mut events = state.stats_events.subscribe(); - async_stream::stream! { - let mut snapshots = tokio::time::interval(Duration::from_secs(1)); - snapshots.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); - let mut snapshot_id = 1_u64; - loop { - if !state.stats_stream_enabled.load(Ordering::Relaxed) { - break; +type ResponseStream = Pin> + Send + 'static>>; + +#[tonic::async_trait] +impl proto::kv_dc_relay_server::KvDcRelay for MockKvDcRelay { + type WatchKvCuckooFilterStream = ResponseStream; + type WatchKvUsageStream = ResponseStream; + type WatchLoadStream = ResponseStream; + + async fn watch_kv_cuckoo_filter( + &self, + _request: Request<()>, + ) -> Result, Status> { + Err(Status::unimplemented( + "mock Dynamo does not model the CKF stream", + )) + } + + async fn watch_kv_usage( + &self, + _request: Request<()>, + ) -> Result, Status> { + ensure_enabled(&self.state)?; + let state = self.state.clone(); + let stream = async_stream::stream! { + let mut interval = snapshot_interval(); + loop { + interval.tick().await; + if !state.stats_stream_enabled.load(Ordering::Relaxed) { + break; + } + yield Ok(usage_snapshot(&state).await); } - tokio::select! { - event = events.recv() => match event { - Ok(event) => yield request_update(event), - Err(broadcast::error::RecvError::Lagged(_)) => break, - Err(broadcast::error::RecvError::Closed) => break, - }, - _ = snapshots.tick() => { - let stats = state.kv_cache.lock().await.stats(&state.model_name); - yield kv_update(snapshot_id, stats); - snapshot_id = snapshot_id.saturating_add(1); + }; + Ok(Response::new(Box::pin(stream))) + } + + async fn watch_load( + &self, + _request: Request<()>, + ) -> Result, Status> { + ensure_enabled(&self.state)?; + let state = self.state.clone(); + let stream = async_stream::stream! { + let mut events = state.stats_events.subscribe(); + let mut accumulator = LoadAccumulator::default(); + let mut interval = snapshot_interval(); + loop { + tokio::select! { + event = events.recv() => match event { + Ok(event) => accumulator.observe(event), + Err(broadcast::error::RecvError::Lagged(_)) => { + yield Err(Status::resource_exhausted("mock load event stream lagged")); + break; + } + Err(broadcast::error::RecvError::Closed) => break, + }, + _ = interval.tick() => { + if !state.stats_stream_enabled.load(Ordering::Relaxed) { + break; + } + yield Ok(accumulator.snapshot(&state.model_name)); + } } } - } + }; + Ok(Response::new(Box::pin(stream))) } } -fn request_update(event: StatsStreamEvent) -> proto::StatsUpdate { - proto::StatsUpdate { - update: Some(proto::stats_update::Update::RequestStats( - proto::RequestStats { - request_id: event.request_id, - model: event.model, - tokens_processed: event.tokens_processed, - tokens_generated: event.tokens_generated, - finished: event.finished, - }, - )), +fn ensure_enabled(state: &AppState) -> Result<(), Status> { + if state.stats_stream_enabled.load(Ordering::Relaxed) { + Ok(()) + } else { + Err(Status::unavailable("mock Relay stats streams disabled")) } } -fn kv_update(snapshot_id: u64, stats: crate::kv_cache::KvCacheStats) -> proto::StatsUpdate { - proto::StatsUpdate { - update: Some(proto::stats_update::Update::KvStats( - proto::KvStatsSnapshot { - snapshot_id, - observed_at_unix_ms: crate::openai::unix_millis(), - models: vec![proto::ModelKvStats { - model: stats.model, - aliases: Vec::new(), - routing_cache: Some(proto::RoutingCacheStats { - role: proto::WorkerRole::Aggregated as i32, - capacity_tokens: stats.kv_cache_capacity_tokens, - used_tokens: stats.kv_cache_used_tokens, - free_tokens: stats.kv_cache_free_tokens, - }), - pools: Vec::new(), - }], - }, - )), - } +fn snapshot_interval() -> tokio::time::Interval { + let mut interval = tokio::time::interval(SNAPSHOT_INTERVAL); + interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip); + interval } -pub(crate) fn ndjson_event(update: proto::StatsUpdate) -> Bytes { - let value = match update - .update - .expect("mock stats update must have a payload") - { - proto::stats_update::Update::RequestStats(event) => serde_json::json!({ - "v": 1, - "type": "stats", - "request_id": event.request_id, - "model": event.model, - "tokens_processed": event.tokens_processed, - "tokens_generated": event.tokens_generated, - "finished": event.finished, - }), - proto::stats_update::Update::KvStats(snapshot) => serde_json::json!({ - "v": 1, - "type": "kv_stats_snapshot", - "snapshot_id": snapshot.snapshot_id, - "observed_at_unix_ms": snapshot.observed_at_unix_ms, - "models": snapshot.models.into_iter().map(|model| serde_json::json!({ - "model": model.model, - "aliases": model.aliases, - "routing_cache": model.routing_cache.map(|cache| serde_json::json!({ - "role": "aggregated", - "capacity_tokens": cache.capacity_tokens, - "used_tokens": cache.used_tokens, - "free_tokens": cache.free_tokens, - })), - "pools": [], - })).collect::>(), - }), - }; - let mut line = serde_json::to_vec(&value).expect("mock stats update should serialize"); - line.push(b'\n'); - Bytes::from(line) +#[derive(Default)] +struct LoadAccumulator { + live: HashMap, + windows: BTreeMap, } -#[derive(Clone)] -struct MockFrontendStats { - state: AppState, +struct LiveRequest { + model: String, + input_tokens: u64, + output_tokens: u64, } -#[tonic::async_trait] -impl proto::frontend_stats_server::FrontendStats for MockFrontendStats { - type WatchStatsStream = - Pin> + Send + 'static>>; - type WatchKvPlacementsStream = - Pin> + Send + 'static>>; +#[derive(Default)] +struct WindowCounters { + requests_started: u64, + requests_completed: u64, + input_tokens: u64, + output_tokens: u64, +} - async fn watch_stats( - &self, - _request: Request, - ) -> Result, Status> { - if !self.state.stats_stream_enabled.load(Ordering::Relaxed) { - return Err(Status::unavailable("mock stats stream disabled")); +impl LoadAccumulator { + fn observe(&mut self, event: StatsStreamEvent) { + let window = self.windows.entry(event.model.clone()).or_default(); + if let Some(request) = self.live.get_mut(&event.request_id) { + window.input_tokens = window + .input_tokens + .saturating_add(event.input_tokens.saturating_sub(request.input_tokens)); + window.output_tokens = window + .output_tokens + .saturating_add(event.output_tokens.saturating_sub(request.output_tokens)); + request.input_tokens = request.input_tokens.max(event.input_tokens); + request.output_tokens = request.output_tokens.max(event.output_tokens); + } else { + window.requests_started = window.requests_started.saturating_add(1); + window.input_tokens = window.input_tokens.saturating_add(event.input_tokens); + window.output_tokens = window.output_tokens.saturating_add(event.output_tokens); + self.live.insert( + event.request_id.clone(), + LiveRequest { + model: event.model.clone(), + input_tokens: event.input_tokens, + output_tokens: event.output_tokens, + }, + ); + } + if event.finished { + self.live.remove(&event.request_id); + window.requests_completed = window.requests_completed.saturating_add(1); } - let stream = futures::StreamExt::map(mixed_stats_stream(self.state.clone()), Ok); - Ok(GrpcResponse::new(Box::pin(stream))) } - async fn watch_kv_placements( - &self, - _request: Request, - ) -> Result, Status> { - Err(Status::unimplemented( - "mock Dynamo does not model KV placements", - )) + fn snapshot(&mut self, configured_model: &str) -> proto::LoadSnapshot { + let mut model_ids = BTreeSet::from([configured_model.to_string()]); + model_ids.extend(self.live.values().map(|request| request.model.clone())); + model_ids.extend(self.windows.keys().cloned()); + let windows = std::mem::take(&mut self.windows); + let models = model_ids + .into_iter() + .map(|model| { + let window = windows.get(&model); + let live = self + .live + .values() + .filter(|request| request.model == model) + .collect::>(); + let input_processing = live + .iter() + .filter(|request| request.output_tokens == 0) + .count() as u64; + let output_generation = live.len() as u64 - input_processing; + let pending_input_tokens = live + .iter() + .filter(|request| request.output_tokens == 0) + .map(|request| request.input_tokens) + .sum(); + let live_input_tokens = live.iter().map(|request| request.input_tokens).sum(); + proto::ModelLoad { + model: Some(model_registration(&model)), + ready_frontends: Some(1), + pending_first_output_requests: Some(input_processing), + pending_first_output_input_tokens: Some(pending_input_tokens), + live_input_tokens: Some(live_input_tokens), + input_processing_requests: Some(input_processing), + output_generation_requests: Some(output_generation), + serving_pools: vec![pool_identity()], + requests_started: window.map_or(0, |window| window.requests_started), + requests_completed: window.map_or(0, |window| window.requests_completed), + requests_failed: 0, + requests_cancelled: 0, + input_tokens: Some(window.map_or(0, |window| window.input_tokens)), + output_tokens: window.map_or(0, |window| window.output_tokens), + status: proto::DataStatus::Complete as i32, + expected_frontends: 1, + observed_frontends: 1, + source_observed_at_unix_ms: crate::openai::unix_millis(), + } + }) + .collect(); + proto::LoadSnapshot { + metadata: Some(metadata()), + window_ms: 1_000, + pools: vec![proto::PoolLoad { + pool: Some(pool_identity()), + role: proto::WorkerRole::Aggregated as i32, + live_workers: Some(1), + active_prefill_tokens: None, + active_decode_blocks: None, + max_concurrency: Some(1), + scheduler_status: proto::DataStatus::Complete as i32, + scheduler_observed_at_unix_ms: crate::openai::unix_millis(), + }], + models, + } + } +} + +async fn usage_snapshot(state: &AppState) -> proto::KvUsageSnapshot { + let stats = state.kv_cache.lock().await.stats(&state.model_name); + proto::KvUsageSnapshot { + metadata: Some(metadata()), + pools: vec![proto::PoolKvUsage { + pool: Some(pool_identity()), + models: vec![model_registration(&state.model_name)], + role: proto::WorkerRole::Aggregated as i32, + block_size_tokens: 1, + expected_ranks: 1, + observed_ranks: 1, + capacity_blocks: Some(stats.kv_cache_capacity_tokens), + used_blocks: Some(stats.kv_cache_used_tokens), + status: proto::DataStatus::Complete as i32, + source_observed_at_unix_ms: crate::openai::unix_millis(), + }], + } +} + +fn metadata() -> proto::RelayMessageMetadata { + proto::RelayMessageMetadata { + drt_instance_id: 1, + relay_incarnation: 1, + observed_at_unix_ms: crate::openai::unix_millis(), + } +} + +fn model_registration(model: &str) -> proto::ModelRegistration { + proto::ModelRegistration { + model: model.to_string(), + base_model: model.to_string(), + adapter: None, + aliases: Vec::new(), + } +} + +fn pool_identity() -> proto::PoolIdentity { + proto::PoolIdentity { + cache_semantics_digest: vec![1; 16], + cache_semantics_source: proto::IdentitySource::DefaultDerived as i32, + routing_scope_digest: vec![2; 16], + routing_scope_source: proto::IdentitySource::DefaultDerived as i32, + dc_id: 1, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn load_snapshots_replace_window_counters_and_keep_live_gauges() { + let mut accumulator = LoadAccumulator::default(); + accumulator.observe(StatsStreamEvent { + request_id: "req-1".to_string(), + model: "model-a".to_string(), + input_tokens: 10, + output_tokens: 2, + finished: false, + }); + + let first = accumulator.snapshot("model-a"); + assert_eq!(first.models[0].requests_started, 1); + assert_eq!(first.models[0].input_tokens, Some(10)); + assert_eq!(first.models[0].output_tokens, 2); + assert_eq!(first.models[0].output_generation_requests, Some(1)); + + let second = accumulator.snapshot("model-a"); + assert_eq!(second.models[0].requests_started, 0); + assert_eq!(second.models[0].input_tokens, Some(0)); + assert_eq!(second.models[0].output_tokens, 0); + assert_eq!(second.models[0].output_generation_requests, Some(1)); } } diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs index 032582602..3154b9f93 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/tests.rs @@ -68,9 +68,10 @@ async fn stats_stream_test_control_does_not_disable_health() { .await; assert_eq!(status, axum::http::StatusCode::NO_CONTENT); - assert_eq!( - stats_stream(State(state.clone())).await.status(), - axum::http::StatusCode::SERVICE_UNAVAILABLE + assert!( + !state + .stats_stream_enabled + .load(std::sync::atomic::Ordering::Relaxed) ); assert_eq!(health(State(state)).await, "ok"); } @@ -675,7 +676,6 @@ async fn streaming_response_delays_first_data_frame_until_ttft() { }; let app = Router::new() .route("/v1/chat/completions", post(chat_completions)) - .route("/v1/stats/stream", get(stats_stream)) .with_state(state); let (addr, server) = spawn_test_app(app).await; @@ -702,87 +702,6 @@ async fn streaming_response_delays_first_data_frame_until_ttft() { server.abort(); } -#[tokio::test] -async fn streaming_response_exposes_stats_stream_endpoint() { - let state = AppState { - num_tokens: 2, - ..test_state() - }; - let app = Router::new() - .route("/v1/chat/completions", post(chat_completions)) - .route("/v1/stats/stream", get(stats_stream)) - .with_state(state); - let (addr, server) = spawn_test_app(app).await; - - let body = r#"{"model":"dummy-model","messages":[{"role":"user","content":"hello world"}],"max_tokens":2,"stream":true}"#; - let mut stream = send_json_request( - addr, - "POST", - "/v1/chat/completions", - "x-request-id: req-contract\r\nx-input-tokens: 11", - body, - ) - .await; - - let response = read_until_done(&mut stream).await; - assert!(!response.contains(r#""usage":"#)); - - let mut stream = tokio::net::TcpStream::connect(addr) - .await - .expect("test client should connect"); - stream - .write_all( - format!("GET /v1/stats/stream HTTP/1.1\r\nhost: {addr}\r\nconnection: close\r\n\r\n",) - .as_bytes(), - ) - .await - .expect("stats stream request should write"); - - let response = tokio::time::timeout( - Duration::from_secs(2), - read_until_contains(&mut stream, "application/x-ndjson"), - ) - .await - .expect("stats stream headers should arrive") - .expect("stats stream response should read"); - assert!(response.starts_with("HTTP/1.1 200 OK")); - assert!(response.contains("content-type: application/x-ndjson")); - server.abort(); -} - -#[test] -fn stats_stream_events_are_ndjson() { - let event = stargate_proto::dynamo_frontend_stats::StatsUpdate { - update: Some( - stargate_proto::dynamo_frontend_stats::stats_update::Update::RequestStats( - stargate_proto::dynamo_frontend_stats::RequestStats { - request_id: "req-1".to_string(), - model: "dummy-model".to_string(), - tokens_processed: Some(11), - tokens_generated: Some(2), - finished: true, - }, - ), - ), - }; - - let line = ndjson_event(event); - let value: serde_json::Value = serde_json::from_slice(&line).unwrap(); - - assert_eq!( - value, - serde_json::json!({ - "v": 1, - "type": "stats", - "request_id": "req-1", - "model": "dummy-model", - "tokens_processed": 11, - "tokens_generated": 2, - "finished": true, - }) - ); -} - #[tokio::test] async fn responses_endpoint_streams_response_events_without_private_stats_headers() { let state = AppState { @@ -894,25 +813,6 @@ async fn read_until_sse_data(stream: &mut tokio::net::TcpStream) -> std::io::Res } } -async fn read_until_done(stream: &mut tokio::net::TcpStream) -> String { - let mut bytes = Vec::new(); - let mut buffer = [0u8; 1024]; - loop { - let read = tokio::time::timeout(Duration::from_secs(2), stream.read(&mut buffer)) - .await - .expect("response should continue") - .expect("response should read"); - if read == 0 { - break; - } - bytes.extend_from_slice(&buffer[..read]); - if String::from_utf8_lossy(&bytes).contains("data: [DONE]") { - break; - } - } - String::from_utf8_lossy(&bytes).to_string() -} - async fn read_to_end(stream: &mut tokio::net::TcpStream) -> String { let mut bytes = Vec::new(); stream @@ -932,23 +832,3 @@ async fn raw_http_request(addr: std::net::SocketAddr, request: &str) -> String { .expect("request should write"); read_to_end(&mut stream).await } - -async fn read_until_contains( - stream: &mut tokio::net::TcpStream, - needle: &str, -) -> std::io::Result { - let mut bytes = Vec::new(); - let mut buffer = [0u8; 1024]; - loop { - let read = stream.read(&mut buffer).await?; - if read == 0 { - break; - } - bytes.extend_from_slice(&buffer[..read]); - let text = String::from_utf8_lossy(&bytes); - if text.contains(needle) { - return Ok(text.to_string()); - } - } - Ok(String::from_utf8_lossy(&bytes).to_string()) -} diff --git a/src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto b/src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto deleted file mode 100644 index cf5165a59..000000000 --- a/src/libraries/rust/stargate/crates/proto/proto/dynamo_frontend_stats.proto +++ /dev/null @@ -1,219 +0,0 @@ -// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. -// All rights reserved. SPDX-License-Identifier: Apache-2.0 -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -syntax = "proto3"; - -package dynamo.frontend.stats.v1; - -service FrontendStats -{ - rpc WatchStats(WatchStatsRequest) returns (stream StatsUpdate); - rpc WatchKvPlacements(WatchKvPlacementsRequest) returns (stream KvPlacementUpdate); -} - -message WatchStatsRequest {} - -message StatsUpdate -{ - oneof update - { - RequestStats request_stats = 1; - KvStatsSnapshot kv_stats = 2; - } -} - -message RequestStats -{ - string request_id = 1; - string model = 2; - optional uint64 tokens_processed = 3; - optional uint64 tokens_generated = 4; - bool finished = 5; -} - -message KvStatsSnapshot -{ - uint64 snapshot_id = 1; - uint64 observed_at_unix_ms = 2; - repeated ModelKvStats models = 3; -} - -message ModelKvStats -{ - string model = 1; - repeated string aliases = 2; - optional RoutingCacheStats routing_cache = 3; - repeated KvPoolStats pools = 4; -} - -enum WorkerRole { - WORKER_ROLE_UNSPECIFIED = 0; - WORKER_ROLE_AGGREGATED = 1; - WORKER_ROLE_PREFILL = 2; - WORKER_ROLE_DECODE = 3; - WORKER_ROLE_ENCODE = 4; -} - -message RoutingCacheStats -{ - WorkerRole role = 1; - uint64 capacity_tokens = 2; - uint64 used_tokens = 3; - uint64 free_tokens = 4; -} - -message KvPoolStats -{ - string namespace = 1; - string component = 2; - string endpoint = 3; - WorkerRole role = 4; - StorageTier storage_tier = 5; - uint32 block_size_tokens = 6; - uint64 expected_ranks = 7; - uint64 observed_ranks = 8; - optional uint64 capacity_blocks = 9; - optional uint64 used_blocks = 10; - optional uint64 free_blocks = 11; - optional uint64 active_decode_blocks = 12; - optional uint64 active_prefill_tokens = 13; - bool complete = 14; -} - -message WatchKvPlacementsRequest {} - -message KvPlacementUpdate -{ - oneof update - { - KvPlacementSnapshotBoundary snapshot_begin = 1; - KvPlacementEvents snapshot_events = 2; - KvPlacementSnapshotBoundary snapshot_end = 3; - KvPlacementEvents events = 4; - KvPlacementSourceError source_error = 5; - } -} - -message KvPlacementSnapshotBoundary -{ - uint64 snapshot_id = 1; - bool complete = 2; - repeated KvPlacementCursor cursors = 3; -} - -message KvPlacementCursor -{ - string model = 1; - string namespace = 2; - string component = 3; - string endpoint = 4; - uint64 cursor = 5; -} - -message KvPlacementSource -{ - string model = 1; - string namespace = 2; - string component = 3; - string endpoint = 4; - uint32 block_size_tokens = 5; -} - -message KvPlacementEvents -{ - optional uint64 snapshot_id = 1; - KvPlacementSource source = 2; - uint64 cursor = 3; - uint32 batch_index = 4; - uint32 batch_count = 5; - repeated RouterEvent events = 6; -} - -message KvPlacementSourceError -{ - KvPlacementSource source = 1; - string reason = 2; -} - -enum StorageTier { - STORAGE_TIER_UNSPECIFIED = 0; - STORAGE_TIER_DEVICE = 1; - STORAGE_TIER_HOST_PINNED = 2; - STORAGE_TIER_DISK = 3; - STORAGE_TIER_EXTERNAL = 4; -} - -message ResidencyDomain -{ - enum Kind { - KIND_MISSING = 0; - KIND_WORKER = 1; - KIND_CACHE_OWNER = 2; - KIND_UNKNOWN = 3; - KIND_INVALID = 4; - } - - Kind kind = 1; - string unknown_value = 2; -} - -message RouterEvent -{ - uint64 worker_id = 1; - StorageTier storage_tier = 2; - ResidencyDomain residency_domain = 3; - optional string state_source = 4; - uint64 event_id = 5; - uint32 dp_rank = 6; - oneof data - { - KvCacheStore stored = 7; - KvCacheRemove removed = 8; - KvCacheClear cleared = 9; - } -} - -message KvCacheStore -{ - optional uint64 parent_hash = 1; - optional uint32 start_position = 2; - repeated KvCacheBlock blocks = 3; -} - -message KvCacheBlock -{ - uint64 block_hash = 1; - uint64 tokens_hash = 2; - repeated MultimodalObject multimodal_objects = 3; -} - -message MultimodalObject -{ - uint64 hash = 1; - repeated TokenRange offsets = 2; -} - -message TokenRange -{ - uint64 start = 1; - uint64 end = 2; -} - -message KvCacheRemove -{ - repeated uint64 block_hashes = 1; -} - -message KvCacheClear {} diff --git a/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto b/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto new file mode 100644 index 000000000..8f66678bb --- /dev/null +++ b/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto @@ -0,0 +1,201 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +syntax = "proto3"; + +package dynamo.kvdc.relay.v1; + +import "google/protobuf/empty.proto"; + +service KvDcRelay +{ + rpc WatchKvCuckooFilter(google.protobuf.Empty) returns (stream KvCuckooFilterUpdate); + rpc WatchKvUsage(google.protobuf.Empty) returns (stream KvUsageSnapshot); + rpc WatchLoad(google.protobuf.Empty) returns (stream LoadSnapshot); +} + +message RelayMessageMetadata +{ + fixed64 drt_instance_id = 1; + fixed64 relay_incarnation = 2; + uint64 observed_at_unix_ms = 3; +} + +enum DataStatus { + DATA_STATUS_UNSPECIFIED = 0; + DATA_STATUS_COMPLETE = 1; + DATA_STATUS_DEGRADED = 2; + DATA_STATUS_UNAVAILABLE = 3; +} + +enum IdentitySource { + IDENTITY_SOURCE_UNSPECIFIED = 0; + IDENTITY_SOURCE_DEFAULT_DERIVED = 1; + IDENTITY_SOURCE_EXPLICIT = 2; +} + +enum WorkerRole { + WORKER_ROLE_UNSPECIFIED = 0; + WORKER_ROLE_AGGREGATED = 1; + WORKER_ROLE_PREFILL = 2; + WORKER_ROLE_DECODE = 3; + WORKER_ROLE_ENCODE = 4; +} + +message PoolIdentity +{ + bytes cache_semantics_digest = 1; + IdentitySource cache_semantics_source = 2; + bytes routing_scope_digest = 3; + IdentitySource routing_scope_source = 4; + fixed64 dc_id = 5; +} + +message ModelRegistration +{ + string model = 1; + string base_model = 2; + optional string adapter = 3; + repeated string aliases = 4; +} + +message CuckooFormat +{ + uint32 format_version = 1; + fixed64 seed = 2; + uint64 bucket_count = 3; + uint32 fingerprint_bits = 4; + uint32 slots_per_bucket = 5; +} + +message CuckooProducerIdentity +{ + PoolIdentity pool = 1; + fixed64 producer_incarnation = 2; + uint64 layout_generation = 3; + CuckooFormat format = 4; +} + +message CuckooPoolSnapshot +{ + CuckooProducerIdentity producer = 1; + uint64 sequence = 2; + bytes packed_buckets = 3; + DataStatus status = 4; + CuckooPoolStats stats = 5; +} + +message CuckooBucketImage +{ + uint64 bucket_index = 1; + fixed64 packed_bucket = 2; +} + +message CuckooPoolDelta +{ + CuckooProducerIdentity producer = 1; + uint64 base_sequence = 2; + uint64 sequence = 3; + repeated CuckooBucketImage buckets = 4; +} + +message CuckooPoolRetired +{ + PoolIdentity pool = 1; + fixed64 producer_incarnation = 2; + uint64 layout_generation = 3; +} + +message CuckooPoolStats +{ + CuckooProducerIdentity producer = 1; + uint64 publication_sequence = 2; + uint64 materialized_ranks = 3; + uint64 unique_blocks = 4; + DataStatus status = 5; + uint64 capacity_omissions = 6; + uint64 source_observed_at_unix_ms = 7; +} + +message CuckooStreamHeartbeat +{ + uint64 catalog_revision = 1; + bool initial_sync_complete = 2; +} + +message KvCuckooFilterUpdate +{ + RelayMessageMetadata metadata = 1; + oneof update + { + CuckooPoolSnapshot snapshot = 2; + CuckooPoolDelta delta = 3; + CuckooPoolRetired retired = 4; + CuckooPoolStats stats = 5; + CuckooStreamHeartbeat heartbeat = 6; + } +} + +message KvUsageSnapshot +{ + RelayMessageMetadata metadata = 1; + repeated PoolKvUsage pools = 2; +} + +message PoolKvUsage +{ + PoolIdentity pool = 1; + repeated ModelRegistration models = 2; + WorkerRole role = 3; + uint32 block_size_tokens = 4; + uint64 expected_ranks = 5; + uint64 observed_ranks = 6; + optional uint64 capacity_blocks = 7; + optional uint64 used_blocks = 8; + DataStatus status = 9; + uint64 source_observed_at_unix_ms = 10; +} + +message LoadSnapshot +{ + RelayMessageMetadata metadata = 1; + uint32 window_ms = 2; + repeated PoolLoad pools = 3; + repeated ModelLoad models = 4; +} + +message PoolLoad +{ + PoolIdentity pool = 1; + WorkerRole role = 2; + optional uint64 live_workers = 3; + optional uint64 active_prefill_tokens = 4; + optional uint64 active_decode_blocks = 5; + optional uint64 max_concurrency = 6; + DataStatus scheduler_status = 7; + uint64 scheduler_observed_at_unix_ms = 8; +} + +message ModelLoad +{ + ModelRegistration model = 1; + optional uint64 ready_frontends = 2; + optional uint64 pending_first_output_requests = 3; + optional uint64 pending_first_output_input_tokens = 4; + optional uint64 live_input_tokens = 5; + optional uint64 input_processing_requests = 6; + optional uint64 output_generation_requests = 7; + repeated PoolIdentity serving_pools = 8; + + uint64 requests_started = 10; + uint64 requests_completed = 11; + uint64 requests_failed = 12; + uint64 requests_cancelled = 13; + optional uint64 input_tokens = 14; + uint64 output_tokens = 15; + + DataStatus status = 20; + uint32 expected_frontends = 21; + uint32 observed_frontends = 22; + uint64 source_observed_at_unix_ms = 23; +} diff --git a/src/libraries/rust/stargate/crates/proto/src/build_plan.rs b/src/libraries/rust/stargate/crates/proto/src/build_plan.rs index 1863c81f6..6d284ba9d 100644 --- a/src/libraries/rust/stargate/crates/proto/src/build_plan.rs +++ b/src/libraries/rust/stargate/crates/proto/src/build_plan.rs @@ -50,7 +50,7 @@ pub(crate) fn proto_compile_plans() -> [ProtoCompilePlan; 3] { field_attributes: &[], }, ProtoCompilePlan { - protos: &["proto/dynamo_frontend_stats.proto"], + protos: &["proto/dynamo_kv_dc_relay.proto"], includes: &["proto"], build_server: true, type_attributes: &[], @@ -85,10 +85,10 @@ mod tests { } #[test] - fn dynamo_frontend_stats_plan_builds_client_and_server_proto() { + fn dynamo_kv_dc_relay_plan_builds_client_and_server_proto() { let [_, _, stats_plan] = proto_compile_plans(); - assert_eq!(stats_plan.protos, ["proto/dynamo_frontend_stats.proto"]); + assert_eq!(stats_plan.protos, ["proto/dynamo_kv_dc_relay.proto"]); assert_eq!(stats_plan.includes, ["proto"]); assert!(stats_plan.build_server); } diff --git a/src/libraries/rust/stargate/crates/proto/src/lib.rs b/src/libraries/rust/stargate/crates/proto/src/lib.rs index 7ba6cf200..b0f64b6a7 100644 --- a/src/libraries/rust/stargate/crates/proto/src/lib.rs +++ b/src/libraries/rust/stargate/crates/proto/src/lib.rs @@ -17,8 +17,8 @@ pub mod gateway_pb { tonic::include_proto!("llm_gateway"); } -pub mod dynamo_frontend_stats { - tonic::include_proto!("dynamo.frontend.stats.v1"); +pub mod dynamo_kv_dc_relay { + tonic::include_proto!("dynamo.kvdc.relay.v1"); } pub const REGISTRATION_HEARTBEAT_MS_METADATA: &str = "x-stargate-registration-heartbeat-ms"; @@ -132,7 +132,7 @@ mod tests { assert!(plans[1].build_server); assert!(plans[1].type_attributes.is_empty()); assert!(plans[1].field_attributes.is_empty()); - assert_eq!(plans[2].protos, ["proto/dynamo_frontend_stats.proto"]); + assert_eq!(plans[2].protos, ["proto/dynamo_kv_dc_relay.proto"]); assert_eq!(plans[2].includes, ["proto"]); assert!(plans[2].build_server); assert_eq!(crate::build_script::planned_proto_compile_count(), 3); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs index a5e3d4a68..0e5492484 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup.rs @@ -928,13 +928,6 @@ mod tests { .get(HEADER_REQUEST_ID) .and_then(|value| value.to_str().ok()) { - assert_eq!( - headers - .get("request-id") - .and_then(|value| value.to_str().ok()), - Some(request_id), - "Pylon-generated requests must use one canonical upstream ID" - ); request_ids.lock().await.push(request_id.to_string()); } let prompt = request diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs index 93007d2e6..29092e754 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs @@ -116,10 +116,12 @@ pub(super) async fn send_completion_request( "/v1/chat/completions", )) .header(HEADER_REQUEST_ID, &request_id) - .header("request-id", &request_id) .header(HEADER_MODEL, model_id) .header(HEADER_INPUT_TOKENS, input_tokens.to_string()) .json(request); + if let Some(observer) = observer.as_mut() { + observer.on_upstream_send(); + } let response = match timeout { Some(timeout) => request.timeout(timeout), None => request, diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/generated_request_id.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/generated_request_id.rs index bf58e861d..e3fd33581 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/generated_request_id.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/generated_request_id.rs @@ -59,6 +59,7 @@ pub(crate) fn next_generated_request_id( ) } +#[cfg(test)] pub(crate) fn generated_request_generation( request_id: &str, model_id: &str, diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs index 0daf85ff6..627655932 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/model_lifecycle.rs @@ -975,8 +975,17 @@ mod tests { .iter() .map(|model_id| (*model_id).to_string()) .collect::>(); - wait_for("advertised model set should converge", || { - runtime_state.advertised_model_ids() == expected + wait_for("active model set should converge", || { + let mut active = runtime_state + .advertised_models() + .into_iter() + .filter_map(|(model_id, registration)| { + (registration.status == InferenceServerStatus::Active as i32) + .then_some(model_id) + }) + .collect::>(); + active.sort_unstable(); + active == expected }) .await; } @@ -1694,7 +1703,11 @@ mod tests { upstream.set_models(&["model-a"]).await; let retired = upstream.next_calibration().await; - assert!(runtime_state.advertised_model_ids().is_empty()); + assert_eq!(runtime_state.advertised_model_ids(), ["model-a"]); + assert_eq!( + runtime_state.advertised_models()["model-a"].status, + InferenceServerStatus::Inactive as i32 + ); upstream.set_models(&[]).await; wait_for_generation_retirement(&runtime_state, "model-a").await; let polls = upstream.discovery_polls.load(Ordering::SeqCst); @@ -1789,11 +1802,19 @@ mod tests { error_observed, "the first failed attempt should be recorded" ); - assert_eq!(runtime_state.advertised_model_ids(), ["model-a"]); + assert_eq!(runtime_state.advertised_model_ids(), ["model-a", "model-b"]); + assert_eq!( + runtime_state.advertised_models()["model-b"].status, + InferenceServerStatus::Inactive as i32 + ); let second_attempt = upstream.next_calibration().await; assert_eq!(second_attempt.model_id(), "model-b"); assert_eq!(first_attempt, second_attempt); - assert_eq!(runtime_state.advertised_model_ids(), ["model-a"]); + assert_eq!(runtime_state.advertised_model_ids(), ["model-a", "model-b"]); + assert_eq!( + runtime_state.advertised_models()["model-b"].status, + InferenceServerStatus::Inactive as i32 + ); wait_for_calibration_count(&metrics, "model-b", "error", 2).await; failing.store(false, Ordering::SeqCst); @@ -2120,7 +2141,11 @@ mod tests { runtime_state.current_generation("model-a"), Some(generation.clone()) ); - assert!(runtime_state.advertised_model_ids().is_empty()); + assert_eq!(runtime_state.advertised_model_ids(), ["model-a"]); + assert_eq!( + runtime_state.advertised_models()["model-a"].status, + InferenceServerStatus::Inactive as i32 + ); wait_for_model_ids(&runtime_state, &["model-a"]).await; assert_eq!( runtime_state.current_generation("model-a"), diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/queue_admission.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/queue_admission.rs index 0a3697abe..22a6341c5 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/queue_admission.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/queue_admission.rs @@ -214,14 +214,6 @@ impl QueueAdmissionDecision { } impl LiveRequestState { - pub(crate) fn request_generation(&self, request_id: &str) -> Option { - self.inner - .lock() - .requests - .get(request_id) - .map(|request| request.generation().clone()) - } - pub(crate) fn update_generation_throughput( &self, generation: &ModelGeneration, @@ -767,7 +759,7 @@ impl TrackedPromptPhase { } impl QueueTrackedRequestGuard { - pub(crate) fn on_upstream_response_headers(&mut self) { + pub(crate) fn on_upstream_send(&mut self) { let mut state = self.live_requests.inner.lock(); state.advance_request_phase(&self.request_id, TrackedPromptPhase::InputProcessing); } @@ -1025,10 +1017,10 @@ mod tests { let _priority_two = live_requests.track_request(&required("req-p2", 2, 20)); let mut priority_four = live_requests.track_request(&required("req-p4", 4, 30)); let mut zero_input = live_requests.track_request(&required("req-zero", 1, 0)); - priority_four.on_upstream_response_headers(); + priority_four.on_upstream_send(); assert_eq!(live_requests.snapshot_model("model-a").queue_size, 3); - zero_input.on_upstream_response_headers(); + zero_input.on_upstream_send(); let snapshot = live_requests.snapshot_model("model-a"); assert_eq!(snapshot.queue_size, 3); @@ -1148,7 +1140,7 @@ mod tests { let live_requests = LiveRequestState::default(); let mut request = live_requests.track_request(&required("req-output", 0, 100)); request.observe_output(); - request.on_upstream_response_headers(); + request.on_upstream_send(); let snapshot = live_requests.snapshot_model("model-a"); assert_eq!(snapshot.queue_size, 0); @@ -1201,7 +1193,7 @@ mod tests { let live_requests = LiveRequestState::default(); let mut request = live_requests.track_request(&required("req-progress", 0, 100)); - request.on_upstream_response_headers(); + request.on_upstream_send(); let snapshot = live_requests.snapshot_model("model-a"); assert_eq!(snapshot.queued_input_size, 100); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs index f3110aad1..070f3faa0 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs @@ -62,19 +62,21 @@ pub const DEFAULT_PRIORITY_CEILING: u32 = 3600; pub(crate) mod dynamo { use reqwest::header::{HeaderMap, HeaderName, HeaderValue}; - use stargate_protocol::tunnel_contract::{HEADER_MODEL, HEADER_REQUEST_ID, HEADER_ROUTING_KEY}; + use stargate_protocol::tunnel_contract::{HEADER_MODEL, HEADER_PRIORITY, HEADER_ROUTING_KEY}; /// Engine priority headers pylon derives; the names stay out of the /// shared tunnel contract because only pylon speaks them. - pub(crate) const HEADER_DYNAMO_REQUEST_ID: &str = "request-id"; pub(crate) const HEADER_REQUEST_PRIORITY: &str = "x-dynamo-request-priority"; pub(crate) const HEADER_REQUEST_STRICT_PRIORITY: &str = "x-dynamo-request-strict-priority"; - /// Denylist of engine headers pylon owns: inbound values are stripped in - /// every backend mode so pylon stays their only writer. - const STRIPPED_REQUEST_HEADERS: [&str; 4] = [ - HEADER_DYNAMO_REQUEST_ID, + /// Platform metadata is consumed by pylon and never becomes part of the + /// Dynamo request. Priority is translated below; request state is local. + const STRIPPED_REQUEST_HEADERS: [&str; 7] = [ + "request-id", "x-dynamo-request-id", + HEADER_MODEL, + HEADER_ROUTING_KEY, + HEADER_PRIORITY, HEADER_REQUEST_PRIORITY, HEADER_REQUEST_STRICT_PRIORITY, ]; @@ -83,27 +85,6 @@ pub(crate) mod dynamo { STRIPPED_REQUEST_HEADERS.contains(&name.as_str()) } - /// Translate the validated platform request ID into Dynamo's canonical ID. - pub(crate) fn apply_request_id(request_id: &str, upstream_headers: &mut HeaderMap) { - for name in [HEADER_DYNAMO_REQUEST_ID, "x-dynamo-request-id"] { - upstream_headers.remove(name); - } - for name in [HEADER_MODEL, HEADER_ROUTING_KEY] { - upstream_headers.remove(name); - } - upstream_headers.insert( - HeaderName::from_static(HEADER_DYNAMO_REQUEST_ID), - HeaderValue::from_str(request_id) - .expect("validated x-request-id should be a valid header value"), - ); - debug_assert_eq!( - upstream_headers - .get(HEADER_REQUEST_ID) - .and_then(|value| value.to_str().ok()), - Some(request_id) - ); - } - /// Map the platform rank (lower wins, absent = unconfigured) to the /// engine value (higher wins, read as seconds of queue head start): /// `max(0, ceiling - rank)`, with absent as the lowest value. The head diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs index 12cd1b277..4403b189f 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs @@ -363,6 +363,15 @@ impl TunnelRequestLifecycle { } } + fn on_upstream_send(&mut self) { + if let Some(queue_request) = self.queue_request.as_mut() { + queue_request.on_upstream_send(); + } + if let Some(observer) = self.observer.as_mut() { + observer.on_upstream_send(); + } + } + async fn relay_sse( &mut self, app: &TunnelServerApp, @@ -534,13 +543,10 @@ async fn relay_upstream_response( &app.inference_server_id, )?; transport.send_response_head(status, response_head).await?; - if let Some(lifecycle) = lifecycle.as_mut() { - if let Some(queue_request) = lifecycle.queue_request.as_mut() { - queue_request.on_upstream_response_headers(); - } - if let Some(observer) = lifecycle.observer.as_mut() { - observer.on_upstream_response_headers(response.headers(), status.as_u16()); - } + if let Some(lifecycle) = lifecycle.as_mut() + && let Some(observer) = lifecycle.observer.as_mut() + { + observer.on_upstream_response_headers(response.headers(), status.as_u16()); } if let Some(lifecycle) = lifecycle.as_mut() && lifecycle @@ -684,7 +690,7 @@ pub(super) async fn forward_tunnel_request( &request_headers, body_bytes, health_request, - lifecycle.as_ref(), + lifecycle.as_mut(), ) .await { @@ -724,9 +730,11 @@ async fn send_upstream_request( request_headers: &HeaderMap, body_bytes: Vec, health_request: bool, - lifecycle: Option<&TunnelRequestLifecycle>, + mut lifecycle: Option<&mut TunnelRequestLifecycle>, ) -> Result { - let priority = lifecycle.and_then(|lifecycle| lifecycle.required.priority); + let priority = lifecycle + .as_ref() + .and_then(|lifecycle| lifecycle.required.priority); let span = if !health_request { let span = tracing::info_span!( "pylon_upstream_http_request", @@ -758,12 +766,6 @@ async fn send_upstream_request( span.record("priority", priority); } if app.upstream_backend == UpstreamBackend::Dynamo { - if let Some(lifecycle) = lifecycle { - backend::dynamo::apply_request_id( - &lifecycle.required.request_id, - &mut upstream_headers, - ); - } let dynamo_priority = backend::dynamo::apply_priority_headers( priority, app.priority_ceiling, @@ -776,6 +778,9 @@ async fn send_upstream_request( let send = async { let request_url = join_base_path(&app.upstream_http_base_url, path_and_query) .map_err(UpstreamRequestError::Build)?; + if let Some(lifecycle) = lifecycle.as_deref_mut() { + lifecycle.on_upstream_send(); + } app.http_client .request(method, request_url) .headers(upstream_headers) diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs index ed7b36206..3aec2e3bc 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs @@ -541,21 +541,23 @@ fn pylon_dynamo_priority_headers_are_always_emitted() { } #[test] -fn pylon_translates_platform_request_id_to_dynamo_request_id() { - let mut headers = HeaderMap::new(); - headers.insert("x-request-id", "gateway-request".parse().unwrap()); - headers.insert("x-model", "gateway-model".parse().unwrap()); - headers.insert("x-routing-key", "gateway-route".parse().unwrap()); - headers.insert("request-id", "spoofed-request".parse().unwrap()); - headers.insert("x-dynamo-request-id", "spoofed-legacy".parse().unwrap()); - - dynamo::apply_request_id("gateway-request", &mut headers); - - assert_eq!(headers["x-request-id"], "gateway-request"); - assert_eq!(headers["request-id"], "gateway-request"); - assert!(!headers.contains_key("x-dynamo-request-id")); - assert!(!headers.contains_key("x-model")); - assert!(!headers.contains_key("x-routing-key")); +fn pylon_consumes_platform_metadata_instead_of_forwarding_it_to_dynamo() { + for name in [ + "request-id", + "x-dynamo-request-id", + "x-model", + "x-routing-key", + "x-priority", + ] { + assert!(dynamo::is_stripped_engine_header( + &HeaderName::from_bytes(name.as_bytes()).unwrap() + )); + } + for name in ["x-request-id", "x-input-tokens"] { + assert!(!dynamo::is_stripped_engine_header( + &HeaderName::from_bytes(name.as_bytes()).unwrap() + )); + } } #[test] @@ -1054,20 +1056,16 @@ async fn http3_direct_tunnel_accepts_responses_request_to_upstream() { let app = Router::new().route( "/v1/responses", post(|req: Request| async move { - let model = req - .headers() - .get("x-model") - .and_then(|value| value.to_str().ok()) - .unwrap_or("missing") - .to_string(); + let platform_model_header = req.headers().contains_key("x-model"); let body = axum::body::to_bytes(req.into_body(), 1024 * 1024) .await .unwrap(); - ( - StatusCode::OK, - [(reqwest::header::CONTENT_TYPE.as_str(), "application/json")], - format!(r#"{{"model":"{model}","body_len":{}}}"#, body.len()), - ) + let request: serde_json::Value = serde_json::from_slice(&body).unwrap(); + Json(serde_json::json!({ + "model": request["model"], + "body_len": body.len(), + "platform_model_header": platform_model_header, + })) }), ); let mut config = test_tunnel_config_for(app).await; @@ -1079,11 +1077,12 @@ async fn http3_direct_tunnel_accepts_responses_request_to_upstream() { headers.insert("x-model", "model-h3".parse().unwrap()); headers.insert("x-input-tokens", "7".parse().unwrap()); headers.insert("content-type", "application/json".parse().unwrap()); + let request_body = br#"{"model":"model-h3","input":"hi","stream":true}"#; let response = send_direct_http3_json_request( tunnel.listen_addr(), "/v1/responses?source=http3", headers, - br#"{"input":"hi","stream":true}"#, + request_body, ) .await; @@ -1095,8 +1094,9 @@ async fn http3_direct_tunnel_accepts_responses_request_to_upstream() { ); assert_eq!( payload.get("body_len").and_then(serde_json::Value::as_u64), - Some(28) + Some(request_body.len() as u64) ); + assert_eq!(payload["platform_model_header"], false); tunnel.shutdown().await; } @@ -1637,16 +1637,17 @@ async fn quic_tunnel_forwards_to_http_backend() { let app = Router::new().route( "/v1/chat/completions", post(|req: Request| async move { - let model = req - .headers() - .get("x-model") - .and_then(|v| v.to_str().ok()) - .unwrap_or("none"); + let platform_model_header = req.headers().contains_key("x-model"); let saw_expected_queue_header = req.headers().contains_key("x-stargate-expected-queue-ms"); let saw_retry_control_header = RETRY_CONTROL_REQUEST_HEADERS .iter() .any(|name| req.headers().contains_key(*name)); + let body = axum::body::to_bytes(req.into_body(), 1024 * 1024) + .await + .unwrap(); + let request: serde_json::Value = serde_json::from_slice(&body).unwrap(); + let model = request["model"].as_str().unwrap(); let mut sse = axum::response::Sse::new(async_stream::stream! { yield Ok::<_, std::convert::Infallible>( Event::default().data(r#"{"object":"chat.completion.chunk","choices":[{"delta":{"content":"ok"}}]}"#) @@ -1666,6 +1667,10 @@ async fn quic_tunnel_forwards_to_http_backend() { HeaderName::from_static("x-saw-retry-control"), HeaderValue::from_str(&saw_retry_control_header.to_string()).unwrap(), ); + sse.headers_mut().insert( + HeaderName::from_static("x-saw-platform-model"), + HeaderValue::from_str(&platform_model_header.to_string()).unwrap(), + ); *sse.status_mut() = StatusCode::OK; sse }), @@ -1682,7 +1687,10 @@ async fn quic_tunnel_forwards_to_http_backend() { headers.insert(name, "spoofed".parse().unwrap()); } tunnel - .send(headers, br#"{"messages":[],"stream":true}"#) + .send( + headers, + br#"{"model":"model-a","messages":[],"stream":true}"#, + ) .await; let response_headers = tunnel.response_head(StatusCode::OK).await; @@ -1703,6 +1711,7 @@ async fn quic_tunnel_forwards_to_http_backend() { "false" ); assert_eq!(response_headers["x-saw-retry-control"], "false"); + assert_eq!(response_headers["x-saw-platform-model"], "false"); let response_text = read_response_text(&mut tunnel.recv).await; let events = parse_test_sse_events(&response_text); @@ -1744,7 +1753,7 @@ fn dynamo_priority_echo_router() -> Router { let dynamo_priority = echo_header("x-dynamo-request-priority"); let dynamo_strict_priority = echo_header("x-dynamo-request-strict-priority"); let x_request_id = echo_header("x-request-id"); - let dynamo_request_id = echo_header(dynamo::HEADER_DYNAMO_REQUEST_ID); + let dynamo_request_id = echo_header("request-id"); let platform_routing_identity_present = ["x-model", "x-routing-key"] .into_iter() .any(|name| req.headers().contains_key(name)); @@ -1817,7 +1826,7 @@ async fn quic_tunnel_translates_dynamo_request_headers() { ); assert_eq!(response_headers["x-echo-dynamo-strict-priority"], "0"); assert_eq!(response_headers["x-echo-request-id"], "req-dynamo-1"); - assert_eq!(response_headers["x-echo-dynamo-request-id"], "req-dynamo-1"); + assert_eq!(response_headers["x-echo-dynamo-request-id"], "absent"); assert_eq!(response_headers["x-saw-platform-routing-identity"], "false"); tunnel.shutdown().await; @@ -2734,17 +2743,19 @@ async fn assert_direct_embeddings_case(case: DirectEmbeddingsCase) { let hits = hits_for_app.clone(); async move { hits.fetch_add(1, Ordering::Relaxed); + let platform_model_header = req.headers().contains_key("x-model"); let path = req .uri() .path_and_query() .map_or_else(|| req.uri().path().to_string(), |value| value.to_string()); - let model = req.headers()["x-model"].to_str().unwrap().to_string(); let body = axum::body::to_bytes(req.into_body(), 1024 * 1024) .await .unwrap(); + let request: serde_json::Value = serde_json::from_slice(&body).unwrap(); Json(serde_json::json!({ "path": path, - "model": model, + "model": request["model"], + "platform_model_header": platform_model_header, "body": String::from_utf8(body.to_vec()).unwrap(), "object": "list", "data": serde_json::from_str::(case.response_data_json) @@ -2774,6 +2785,7 @@ async fn assert_direct_embeddings_case(case: DirectEmbeddingsCase) { let payload: serde_json::Value = serde_json::from_slice(&response.body).unwrap(); assert_eq!(payload["path"], case.success_path); assert_eq!(payload["model"], "model-embed"); + assert_eq!(payload["platform_model_header"], false); assert_eq!( payload["body"], String::from_utf8(case.request_body.to_vec()).unwrap() diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer.rs index 84c1192b3..db212e718 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer.rs @@ -49,7 +49,8 @@ pub enum RequestObservationEndpoint { #[derive(Debug)] enum RequestLifecycleState { - UpstreamConnecting, + Queued, + UpstreamSent, Responding(ResponsePhaseData), Terminal { outcome: RequestTerminalOutcome, @@ -77,7 +78,8 @@ impl RequestTerminalOutcome { impl RequestLifecycleState { fn observation_state(&self) -> RequestObservationState { match self { - Self::UpstreamConnecting => RequestObservationState::UpstreamConnecting, + Self::Queued => RequestObservationState::Queued, + Self::UpstreamSent => RequestObservationState::InputProcessing, Self::Responding(response) if response.first_output_at.is_none() => { RequestObservationState::InputProcessing } @@ -90,7 +92,7 @@ impl RequestLifecycleState { match self { Self::Responding(response) => Some(response), Self::Terminal { response, .. } => response.as_ref(), - Self::UpstreamConnecting => None, + Self::Queued | Self::UpstreamSent => None, } } } @@ -180,7 +182,7 @@ impl RequestObserver { input_tokens, generation, embedding_items: None, - state: RequestLifecycleState::UpstreamConnecting, + state: RequestLifecycleState::Queued, runtime_state, }; observer.emit(); @@ -195,14 +197,33 @@ impl RequestObserver { } } + pub(crate) fn on_upstream_send(&mut self) { + match self.state { + RequestLifecycleState::Queued => {} + RequestLifecycleState::UpstreamSent + | RequestLifecycleState::Responding(_) + | RequestLifecycleState::Terminal { .. } => { + panic!( + "invalid upstream-send transition for request_id={} from state={:?}", + self.request_id, + self.state.observation_state() + ) + } + } + self.state = RequestLifecycleState::UpstreamSent; + self.emit(); + } + pub(crate) fn on_upstream_response_headers( &mut self, _response_headers: &HeaderMap, status: u16, ) { match self.state { - RequestLifecycleState::UpstreamConnecting => {} - RequestLifecycleState::Responding(_) | RequestLifecycleState::Terminal { .. } => { + RequestLifecycleState::UpstreamSent => {} + RequestLifecycleState::Queued + | RequestLifecycleState::Responding(_) + | RequestLifecycleState::Terminal { .. } => { panic!( "invalid response-header transition for request_id={} from state={:?}", self.request_id, @@ -327,7 +348,9 @@ impl RequestObserver { "invalid finish transition for request_id={} from state={outcome:?}", self.request_id, ), - RequestLifecycleState::UpstreamConnecting => (RequestTerminalOutcome::Failed, None), + RequestLifecycleState::Queued | RequestLifecycleState::UpstreamSent => { + (RequestTerminalOutcome::Failed, None) + } }; self.state = RequestLifecycleState::Terminal { outcome, response }; @@ -345,7 +368,7 @@ impl RequestObserver { fn terminate(&mut self, outcome: RequestTerminalOutcome, action: &'static str) { let response = match &self.state { RequestLifecycleState::Responding(response) => Some(*response), - RequestLifecycleState::UpstreamConnecting => None, + RequestLifecycleState::Queued | RequestLifecycleState::UpstreamSent => None, RequestLifecycleState::Terminal { outcome: prior, .. } => panic!( "invalid {action} transition for request_id={} from state={prior:?}", self.request_id, @@ -579,6 +602,8 @@ mod tests { let (runtime_state, rx) = observed_runtime(8); let mut observer = test_observer(request_id, runtime_state); recv_observation(&rx, "initial observation should be emitted").await; + observer.on_upstream_send(); + recv_observation(&rx, "upstream-send observation should be emitted").await; observer.on_upstream_response_headers(&HeaderMap::new(), 200); let headers = recv_observation(&rx, "response-header observation should be emitted").await; (observer, rx, headers) @@ -707,10 +732,7 @@ mod tests { .collect::>(); assert_eq!(observations.len(), 2); assert_eq!(observations[0].endpoint, endpoint); - assert_eq!( - observations[0].state, - RequestObservationState::UpstreamConnecting - ); + assert_eq!(observations[0].state, RequestObservationState::Queued); assert_eq!(observations[1].endpoint, endpoint); assert_eq!(observations[1].state, RequestObservationState::Failed); } @@ -730,6 +752,7 @@ mod tests { None, runtime_state, ); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); if let Some(generation) = observer.generation_mut() { generation.observe_output_message(); @@ -778,6 +801,7 @@ mod tests { let (runtime_state, rx) = observed_runtime(8); let mut observer = RequestObserver::accepted(embeddings_required_headers(), runtime_state); observer.update_embedding_items(Some(1)); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); observer.finish(); while rx.try_recv().is_ok() {} @@ -803,6 +827,7 @@ mod tests { #[tokio::test] async fn counts_sse_events_across_chunk_boundaries() { let mut observer = test_observer("req-1", PylonRuntimeState::default()); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); observer.observe_output_message(); observer.observe_output_message(); @@ -823,9 +848,14 @@ mod tests { let mut observer = test_observer("req-live", runtime_state); let initial = recv_observation(&rx, "initial observation should be emitted").await; - assert_eq!(initial.state, RequestObservationState::UpstreamConnecting); + assert_eq!(initial.state, RequestObservationState::Queued); assert!(!initial.is_terminal()); + observer.on_upstream_send(); + let sent = recv_observation(&rx, "upstream-send observation should be emitted").await; + assert_eq!(sent.state, RequestObservationState::InputProcessing); + assert!(!sent.is_terminal()); + observer.on_upstream_response_headers(&HeaderMap::new(), 200); let first = recv_observation(&rx, "response-header observation should be emitted").await; assert_eq!(first.state, RequestObservationState::InputProcessing); @@ -839,16 +869,13 @@ mod tests { } #[tokio::test] - async fn upstream_connecting_observation_is_emitted_when_request_starts() { + async fn queued_observation_is_emitted_when_request_starts() { let (runtime_state, rx) = observed_runtime(8); let _observer = test_observer("req-connect", runtime_state); let observation = recv_observation(&rx, "initial observation should be emitted").await; assert_eq!(observation.request_id, "req-connect"); - assert_eq!( - observation.state, - RequestObservationState::UpstreamConnecting - ); + assert_eq!(observation.state, RequestObservationState::Queued); assert_eq!(observation.upstream_status, None); assert!(!observation.is_terminal()); } @@ -858,7 +885,7 @@ mod tests { let (runtime_state, rx) = observed_runtime(8); let observer = test_observer("req-cancel", runtime_state); let initial = recv_observation(&rx, "initial observation should be emitted").await; - assert_eq!(initial.state, RequestObservationState::UpstreamConnecting); + assert_eq!(initial.state, RequestObservationState::Queued); drop(observer); @@ -870,6 +897,7 @@ mod tests { #[tokio::test] async fn accumulates_output_tokens() { let mut observer = test_observer("req-1", PylonRuntimeState::default()); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); observer.observe_output_message(); observer.observe_output_tokens(3); @@ -1092,6 +1120,7 @@ mod tests { #[test] fn finish_without_output_panics() { let mut observer = test_observer("req-2", PylonRuntimeState::default()); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| observer.finish())); assert!(panic.is_err()); @@ -1110,6 +1139,7 @@ mod tests { #[should_panic(expected = "invalid fail transition")] fn fail_after_complete_panics() { let mut observer = make_test_observer(); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); observer.observe_output_message(); observer.finish(); @@ -1148,6 +1178,7 @@ mod tests { assert_terminal_response_header_panic( |observer| { + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 200); observer.observe_output_message(); observer.finish(); @@ -1161,6 +1192,7 @@ mod tests { #[tokio::test] async fn failed_response_stays_failed() { let mut observer = test_observer("req-3", PylonRuntimeState::default()); + observer.on_upstream_send(); observer.on_upstream_response_headers(&HeaderMap::new(), 503); observer.finish(); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer/tunnel.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer/tunnel.rs index 7f20ffd7d..7ddb186d7 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer/tunnel.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/request_observer/tunnel.rs @@ -51,6 +51,10 @@ impl TunnelRequestObserver { } } + pub(crate) fn on_upstream_send(&mut self) { + self.observer.on_upstream_send(); + } + pub(crate) fn on_upstream_response_headers( &mut self, response_headers: &HeaderMap, diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs index 9070b065f..63b45b1e8 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs @@ -13,7 +13,7 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::collections::HashMap; +use std::collections::{BTreeMap, HashMap}; use std::sync::Arc; use parking_lot::Mutex; @@ -90,7 +90,7 @@ impl ModelGeneration { } } -#[derive(Clone, Debug, Default)] +#[derive(Clone, Debug)] pub struct PylonRuntimeState { advertised: Arc>, live_requests: LiveRequestState, @@ -116,7 +116,9 @@ pub(crate) enum RequestGenerationAdmission { struct AdvertisedRuntimeState { base_status: InferenceServerStatus, require_admitted_generation: bool, + require_relay_load: bool, models: HashMap, + relay_models: BTreeMap, } impl AdvertisedRuntimeState { @@ -138,6 +140,7 @@ struct RuntimeModelState { generation: u64, stats: CurrentModelStats, publication: ModelPublication, + relay_ready: bool, } #[derive(Debug, Default)] @@ -174,7 +177,9 @@ impl PylonRuntimeState { advertised: Arc::new(Mutex::new(AdvertisedRuntimeState { base_status: initial_status, require_admitted_generation: true, + require_relay_load: false, models, + relay_models: BTreeMap::new(), })), live_requests: LiveRequestState::default(), metrics: None, @@ -199,6 +204,28 @@ impl PylonRuntimeState { self.advertised.lock().base_status = status; } + pub(crate) fn require_relay_load(&self, required: bool) { + self.advertised.lock().require_relay_load = required; + } + + pub(crate) fn replace_relay_models(&self, models: BTreeMap) { + let mut advertised = self.advertised.lock(); + for (model_id, model) in &mut advertised.models { + model.relay_ready = models.get(model_id).copied().unwrap_or(false); + } + advertised.relay_models = models; + } + + pub(crate) fn mark_relay_load_unavailable(&self) { + let mut advertised = self.advertised.lock(); + for model in advertised.models.values_mut() { + model.relay_ready = false; + } + for ready in advertised.relay_models.values_mut() { + *ready = false; + } + } + pub(crate) fn model_ids(&self) -> Vec { let mut model_ids = self .advertised @@ -217,10 +244,16 @@ impl PylonRuntimeState { if advertised.models.contains_key(generation.model_id()) { return false; } + let relay_ready = advertised + .relay_models + .get(generation.model_id()) + .copied() + .unwrap_or(false); advertised.models.insert( generation.model_id.clone(), RuntimeModelState { generation: generation.sequence, + relay_ready, ..RuntimeModelState::default() }, ); @@ -241,7 +274,10 @@ impl PylonRuntimeState { ) -> RequestGenerationAdmission { let advertised = self.advertised.lock(); match advertised.models.get(model_id) { - Some(model) if matches!(model.publication, ModelPublication::Admitted { .. }) => { + Some(model) + if matches!(model.publication, ModelPublication::Admitted { .. }) + && (!advertised.require_relay_load || model.relay_ready) => + { RequestGenerationAdmission::Admitted(ModelGeneration::new( model_id, model.generation, @@ -364,12 +400,13 @@ impl PylonRuntimeState { pub(crate) fn advertised_models(&self) -> HashMap { let advertised = self.advertised.lock(); - advertised + let mut registrations = advertised .models .iter() - .filter_map(|(model_id, model)| { - let ModelPublication::Admitted { bringup_ready } = model.publication else { - return None; + .map(|(model_id, model)| { + let bringup_ready = match model.publication { + ModelPublication::Pending => false, + ModelPublication::Admitted { bringup_ready } => bringup_ready, }; let stats = &model.stats; let registration = InferenceServerModelRegistration { @@ -404,23 +441,35 @@ impl PylonRuntimeState { stats_capabilities: stats.stats_capabilities.clone(), stats_sources: stats.stats_sources.clone(), }), - status: gated_model_status(advertised.base_status, bringup_ready).into(), + status: gated_model_status( + advertised.base_status, + bringup_ready && (!advertised.require_relay_load || model.relay_ready), + ) + .into(), }; - Some((model_id.clone(), registration)) + (model_id.clone(), registration) }) - .collect() + .collect::>(); + for model_id in advertised.relay_models.keys() { + registrations.entry(model_id.clone()).or_insert_with(|| { + InferenceServerModelRegistration { + stats: Some(ModelStats::default()), + status: InferenceServerStatus::Inactive.into(), + } + }); + } + registrations } pub fn advertised_model_ids(&self) -> Vec { let advertised = self.advertised.lock(); - let mut model_ids = advertised + let model_ids = advertised .models - .iter() - .filter(|(_, model)| matches!(model.publication, ModelPublication::Admitted { .. })) - .map(|(model_id, _)| model_id.clone()) - .collect::>(); - model_ids.sort_unstable(); - model_ids + .keys() + .chain(advertised.relay_models.keys()) + .cloned() + .collect::>(); + model_ids.into_iter().collect() } pub fn observe_request(&self, observation: RequestObservation) { @@ -520,10 +569,6 @@ impl PylonRuntimeState { .update_active_output_tps(request_id, active_chat_output_tps) } - pub(crate) fn request_generation(&self, request_id: &str) -> Option { - self.live_requests.request_generation(request_id) - } - pub(crate) fn snapshot_live_model(&self, model_id: &str) -> QueueModelSnapshot { self.current_generation(model_id) .map_or_else(QueueModelSnapshot::default, |generation| { @@ -578,6 +623,12 @@ impl PylonRuntimeState { } } +impl Default for PylonRuntimeState { + fn default() -> Self { + Self::new(InferenceServerStatus::Unknown, &[]) + } +} + impl RequestObservationEvent { pub fn observation(&self) -> &RequestObservation { &self.observation @@ -601,6 +652,7 @@ pub(crate) fn gated_model_status( #[cfg(test)] mod tests { + use std::collections::BTreeMap; use std::time::Duration; use stargate_proto::pb::InferenceServerStatus; @@ -726,7 +778,7 @@ mod tests { } #[test] - fn pending_generation_is_omitted_until_exact_publication() { + fn pending_generation_is_advertised_inactive_until_exact_publication() { let runtime_state = PylonRuntimeState::new(InferenceServerStatus::Active, &[]); let generation = ModelGeneration::new("model-a", 1); @@ -739,7 +791,10 @@ mod tests { runtime_state.request_generation_admission("model-a"), RequestGenerationAdmission::Unavailable ); - assert!(runtime_state.advertised_models().is_empty()); + assert_eq!( + runtime_state.advertised_models()["model-a"].status, + InferenceServerStatus::Inactive as i32 + ); assert!(runtime_state.publish_generation(&generation)); assert_eq!( @@ -754,6 +809,59 @@ mod tests { .collect::>(), ["model-a"] ); + assert_eq!( + runtime_state.advertised_models()["model-a"].status, + InferenceServerStatus::Active as i32 + ); + } + + #[test] + fn frontend_and_relay_model_union_is_advertised_with_intersection_active() { + let runtime_state = PylonRuntimeState::new(InferenceServerStatus::Active, &[]); + runtime_state.require_relay_load(true); + runtime_state.replace_relay_models(BTreeMap::from([("relay-only".to_string(), true)])); + + assert_eq!(runtime_state.current_generation("relay-only"), None); + assert_eq!( + runtime_state.advertised_models()["relay-only"].status, + InferenceServerStatus::Inactive as i32 + ); + + let relay_generation = ModelGeneration::new("relay-only", 1); + assert!(runtime_state.begin_generation(relay_generation.clone())); + assert!(runtime_state.publish_generation(&relay_generation)); + assert_eq!( + runtime_state.advertised_models()["relay-only"].status, + InferenceServerStatus::Active as i32 + ); + + assert!(runtime_state.retire_generation(&relay_generation).is_some()); + assert_eq!(runtime_state.current_generation("relay-only"), None); + assert_eq!( + runtime_state.advertised_models()["relay-only"].status, + InferenceServerStatus::Inactive as i32 + ); + + let frontend_generation = ModelGeneration::new("frontend-only", 2); + assert!(runtime_state.begin_generation(frontend_generation.clone())); + assert!(runtime_state.publish_generation(&frontend_generation)); + assert_eq!( + runtime_state.advertised_models()["frontend-only"].status, + InferenceServerStatus::Inactive as i32 + ); + + runtime_state.replace_relay_models(BTreeMap::from([ + ("frontend-only".to_string(), true), + ("relay-only".to_string(), true), + ])); + assert_eq!( + runtime_state.advertised_models()["frontend-only"].status, + InferenceServerStatus::Active as i32 + ); + assert_eq!( + runtime_state.advertised_model_ids(), + ["frontend-only", "relay-only"] + ); } #[test] @@ -777,7 +885,10 @@ mod tests { runtime_state.model_stats("model-a"), Some(super::CurrentModelStats::default()) ); - assert!(runtime_state.advertised_models().is_empty()); + assert_eq!( + runtime_state.advertised_models()["model-a"].status, + InferenceServerStatus::Inactive as i32 + ); } #[test] diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs index 75e67d7ae..b5305685d 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/aggregator.rs @@ -40,12 +40,13 @@ pub(super) struct ModelMetricsState { pub(super) max_chat_output_tps: f64, pub(super) max_embedding_item_tps: f64, pub(super) kv_cache: Option, + pub(super) relay_load: Option, pub(super) input_tps_distribution: TpsDistribution, aggregate_state_counted: bool, pub(super) counter_output_tps_authoritative: bool, pub(super) chunk_usage_stats_observed: bool, pub(super) kv_cache_stats_observed: bool, - pub(super) kv_cache_received_at: Option, + pub(super) relay_load_stats_observed: bool, pub(super) engine_stream_stats_observed: bool, pub(super) last_stats_event_at: Option, pub(super) stats_observed_at_unix_ms: u64, @@ -72,6 +73,27 @@ pub struct KvCacheStatsSnapshot { pub struct KvCacheStatsEnvelope { pub(crate) models: Vec, } + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct RelayLoadStatsSnapshot { + pub(crate) model: String, + pub(crate) input_tps: Option, + pub(crate) output_tps: f64, + pub(crate) queue_size: u64, + pub(crate) queued_input_size: Option, + pub(crate) num_running_queries: u64, + pub(crate) max_engine_concurrency: Option, + pub(crate) total_query_input_size: Option, + pub(crate) input_processing_queries: u64, + pub(crate) output_generation_queries: u64, + pub(crate) source_observed_at_unix_ms: u64, + pub(crate) complete: bool, +} + +#[derive(Debug, Clone, Default, PartialEq)] +pub struct RelayLoadStatsEnvelope { + pub(crate) models: Vec, +} struct RequestCounterState { generation: ModelGeneration, input: CounterSampleState, @@ -331,23 +353,18 @@ impl StatsAggregator { StatsAggregatorUpdate::RequestCounters(update) => { self.apply_request_counters_into(update, updated_models) } - StatsAggregatorUpdate::KvCache(snapshot) => updated_models - .extend(self.apply_kv_cache_snapshot(snapshot, tokio::time::Instant::now())), + StatsAggregatorUpdate::KvCache(snapshot) => { + updated_models.extend(self.apply_kv_cache_snapshot(snapshot)) + } + StatsAggregatorUpdate::RelayLoad(snapshot) => { + updated_models.extend(self.apply_relay_load_snapshot(snapshot)) + } StatsAggregatorUpdate::FinalizeRequest(update) => { updated_models.extend(self.finalize_request(update)) } - StatsAggregatorUpdate::EnableOpenAiFallback => self.enable_openai_fallback(), } } - pub(super) fn apply_control_update(&mut self, update: &StatsAggregatorUpdate) -> bool { - if !matches!(update, StatsAggregatorUpdate::EnableOpenAiFallback) { - return false; - } - self.enable_openai_fallback(); - true - } - pub(super) fn openai_fallback_stats_enabled(&self) -> bool { self.config.openai_fallback_stats_enabled } @@ -467,27 +484,6 @@ impl StatsAggregator { } } - let kv_cache_ttl = self.config.kv_cache_stats_ttl; - if !kv_cache_ttl.is_zero() { - for (model_id, generation_state) in &mut self.per_model { - let state = &mut generation_state.metrics; - if state.kv_cache_received_at.is_some_and(|received_at| { - now.saturating_duration_since(received_at) >= kv_cache_ttl - }) { - state.kv_cache_received_at = None; - if state.kv_cache.take().is_some() { - state.stats_observed_at_unix_ms = current_unix_millis(); - tracing::warn!( - model_id, - ttl_ms = kv_cache_ttl.as_millis(), - "clearing stale KV cache stats" - ); - push_dirty_model(&mut dirty_models, model_id.clone()); - } - } - } - } - if let Some(metrics) = self.runtime_state.metrics() { metrics .observe_engine_stats_model_states(ENGINE_STATS_SOURCE, self.model_state_count()); @@ -521,6 +517,49 @@ impl StatsAggregator { self.publish_request_counter_samples(update, samples, updated_models); } + fn apply_relay_load_snapshot( + &mut self, + snapshot: RelayLoadStatsEnvelope, + ) -> Vec { + let by_model = snapshot + .models + .into_iter() + .map(|model| (model.model.clone(), model)) + .collect::>(); + let mut dirty = Vec::new(); + for (model_id, generation) in &mut self.per_model { + let next = by_model.get(model_id).cloned(); + generation.metrics.relay_load_stats_observed |= next.is_some(); + let input_tps = next + .as_ref() + .filter(|load| load.complete) + .and_then(|load| load.input_tps) + .filter(|input_tps| valid_last_mean_input_tps(*input_tps)); + let input_tps_changed = generation.pinned_input_tps.is_none() + && input_tps.is_some_and(|input_tps| { + if generation.metrics.last_mean_input_tps == input_tps { + false + } else { + generation.metrics.last_mean_input_tps = input_tps; + true + } + }); + if generation.metrics.relay_load != next { + generation.metrics.relay_load = next; + if let Some(load) = &generation.metrics.relay_load { + generation.metrics.stats_observed_at_unix_ms = load.source_observed_at_unix_ms; + } + dirty.push(model_id.clone()); + } else if input_tps_changed { + dirty.push(model_id.clone()); + } + } + dirty + .into_iter() + .map(|model_id| self.snapshot_update(model_id)) + .collect() + } + fn prepare_request_counter_update(&mut self, update: &mut RequestCounterUpdate) -> bool { let Some(generation) = self.resolve_request_generation(&update.request_id, update.generation.as_ref()) @@ -790,21 +829,6 @@ impl StatsAggregator { true } - fn enable_openai_fallback(&mut self) { - if self.config.openai_fallback_stats_enabled { - return; - } - self.config.openai_fallback_stats_enabled = true; - tracing::warn!("OpenAI fallback stats enabled after engine stats stream was unsupported"); - if let Some(metrics) = self.runtime_state.metrics() { - metrics.observe_engine_stats_source_transition( - ENGINE_STATS_SOURCE, - "openai_fallback", - "unsupported", - ); - } - } - fn snapshot_update(&self, model_id: String) -> ModelStatsUpdate { let generation = self .current_generation(&model_id) @@ -907,20 +931,28 @@ impl ModelMetricsState { } else { inputs.active_chat_output_tps }; + let relay = self.relay_load.as_ref().filter(|load| load.complete); CurrentModelStats { last_mean_input_tps: self.last_mean_input_tps, - output_tps: active_chat_output_tps.max(average_with_sum( - &self.chat_output_tps_samples, - self.chat_output_tps_sum, - )), + output_tps: relay.map_or_else( + || { + active_chat_output_tps.max(average_with_sum( + &self.chat_output_tps_samples, + self.chat_output_tps_sum, + )) + }, + |load| load.output_tps, + ), embedding_item_tps: average_with_sum( &self.embedding_item_tps_samples, self.embedding_item_tps_sum, ), max_output_tps: self.max_chat_output_tps, max_embedding_item_tps: self.max_embedding_item_tps, - queue_size: inputs.queue_size, - queued_input_size: inputs.queued_input_size, + queue_size: relay.map_or(inputs.queue_size, |load| load.queue_size), + queued_input_size: relay + .and_then(|load| load.queued_input_size) + .unwrap_or(inputs.queued_input_size), kv_cache_capacity_tokens: kv_cache .map(|snapshot| snapshot.kv_cache_capacity_tokens) .unwrap_or_default(), @@ -936,12 +968,19 @@ impl ModelMetricsState { free_tokens: snapshot.kv_cache_free_tokens, source_observed_at_unix_ms: snapshot.source_observed_at_unix_ms, }), - num_running_queries: inputs.num_running_queries, - max_engine_concurrency: None, - total_query_input_size: inputs.total_query_input_size, + num_running_queries: relay + .map_or(inputs.num_running_queries, |load| load.num_running_queries), + max_engine_concurrency: relay.and_then(|load| load.max_engine_concurrency), + total_query_input_size: relay + .and_then(|load| load.total_query_input_size) + .unwrap_or(inputs.total_query_input_size), queue_time_estimate_ms_by_priority: None, - input_processing_queries: inputs.input_processing_queries, - output_generation_queries: inputs.output_generation_queries, + input_processing_queries: relay.map_or(inputs.input_processing_queries, |load| { + load.input_processing_queries + }), + output_generation_queries: relay.map_or(inputs.output_generation_queries, |load| { + load.output_generation_queries + }), stats_observed_at_unix_ms: self.stats_observed_at_unix_ms, stats_capabilities, stats_sources, @@ -956,7 +995,9 @@ impl ModelMetricsState { self.engine_stream_stats_observed .then_some(("model.throughput.engine_stream", ENGINE_STATS_SOURCE)), self.kv_cache_stats_observed - .then_some(("machine.kv_cache.http", "kv_cache_stats")), + .then_some(("machine.kv_cache.dynamo_relay", "dynamo_relay_kv_usage")), + self.relay_load_stats_observed + .then_some(("model.load.dynamo_relay", "dynamo_relay_load")), ] .into_iter() .flatten() diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs index 8520c8c70..08b1d740e 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/collector.rs @@ -31,7 +31,6 @@ const DEFAULT_SMOOTHING_WINDOW_SIZE: usize = 8; const DEFAULT_MIN_INPUT_TOKENS: u64 = 1; const DEFAULT_MIN_OUTPUT_TOKENS: u64 = 1; const DEFAULT_DURATION_FLOOR: Duration = Duration::from_millis(10); -const DEFAULT_KV_CACHE_STATS_TTL: Duration = Duration::from_secs(5); const DEFAULT_ENGINE_STATS_REQUEST_TTL: Duration = Duration::from_secs(300); const DEFAULT_ENGINE_STATS_MODEL_TTL: Duration = Duration::from_secs(30); const DEFAULT_ENGINE_STATS_SWEEP_INTERVAL: Duration = Duration::from_secs(1); @@ -43,7 +42,6 @@ pub struct StatsCollectorConfig { pub min_input_tokens: u64, pub min_output_tokens: u64, pub duration_floor: Duration, - pub kv_cache_stats_ttl: Duration, pub engine_stats_request_ttl: Duration, pub engine_stats_model_ttl: Duration, pub engine_stats_sweep_interval: Duration, @@ -58,7 +56,6 @@ impl Default for StatsCollectorConfig { min_input_tokens: DEFAULT_MIN_INPUT_TOKENS, min_output_tokens: DEFAULT_MIN_OUTPUT_TOKENS, duration_floor: DEFAULT_DURATION_FLOOR, - kv_cache_stats_ttl: DEFAULT_KV_CACHE_STATS_TTL, engine_stats_request_ttl: DEFAULT_ENGINE_STATS_REQUEST_TTL, engine_stats_model_ttl: DEFAULT_ENGINE_STATS_MODEL_TTL, engine_stats_sweep_interval: DEFAULT_ENGINE_STATS_SWEEP_INTERVAL, @@ -191,8 +188,8 @@ pub enum StatsUpdateSource { pub enum StatsAggregatorUpdate { RequestCounters(RequestCounterUpdate), KvCache(KvCacheStatsEnvelope), + RelayLoad(super::aggregator::RelayLoadStatsEnvelope), FinalizeRequest(FinalizeRequestUpdate), - EnableOpenAiFallback, } #[derive(Debug, Clone)] @@ -270,8 +267,7 @@ pub fn start_stats_collector_with_engine_stats( stats_update_rx: Option>, runtime_state: PylonRuntimeState, ) -> StatsCollectorHandle { - // A wired engine stats stream is the throughput source of truth. Auto mode - // falls back only after the stream task sends EnableOpenAiFallback. + // A wired Relay stream is the aggregate load source of truth. config.openai_fallback_stats_enabled &= stats_update_rx.is_none(); let mut aggregator = StatsAggregator::new(config, runtime_state.clone()); for model_id in runtime_state.model_ids() { @@ -373,9 +369,6 @@ async fn run_stats_collector( stats_update_rx = None; continue; }; - if aggregator.apply_control_update(&update) { - continue; - } stats_aggregator_updated_models.clear(); aggregator.apply_update_into(update, &mut stats_aggregator_updated_models); if let Some(rx) = &stats_update_rx { @@ -435,9 +428,7 @@ fn drain_stats_updates( latest_by_model: &mut IndexMap, ) { drain_ready(updates, |update| { - if !aggregator.apply_control_update(&update) { - aggregator.apply_update_into(update, updated_models); - } + aggregator.apply_update_into(update, updated_models); }); retain_latest_model_updates(updated_models, latest_by_model); } @@ -505,7 +496,10 @@ mod tests { use std::collections::HashMap; use std::sync::Arc; - use super::super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot, StatsAggregator}; + use super::super::aggregator::{ + KvCacheStatsEnvelope, KvCacheStatsSnapshot, RelayLoadStatsEnvelope, RelayLoadStatsSnapshot, + StatsAggregator, + }; use super::super::metrics::PylonMetrics; use super::super::projection::fallback_update_from_observation; use super::*; @@ -977,6 +971,21 @@ mod tests { } } + fn relay_load( + input_tps: Option, + source_observed_at_unix_ms: u64, + ) -> RelayLoadStatsEnvelope { + RelayLoadStatsEnvelope { + models: vec![RelayLoadStatsSnapshot { + model: "model-a".to_string(), + input_tps, + source_observed_at_unix_ms, + complete: true, + ..RelayLoadStatsSnapshot::default() + }], + } + } + fn published_stats( updates: Vec, ) -> CurrentModelStats { @@ -1431,7 +1440,42 @@ mod tests { let mut aggregator = test_aggregator(StatsCollectorConfig::default()); aggregator.apply_kv_cache_stats(kv_cache_stats("model-a")); let stats = aggregator.stream_stats("req-stream-kv", (0, 10), true, Duration::ZERO); - assert_stats!(stats; kv_cache_capacity_tokens: 1_000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600, stats_capabilities: ["model.throughput.engine_stream", "machine.kv_cache.http"], stats_sources: ["engine_stats_stream", "kv_cache_stats"]); + assert_stats!(stats; kv_cache_capacity_tokens: 1_000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600, stats_capabilities: ["model.throughput.engine_stream", "machine.kv_cache.dynamo_relay"], stats_sources: ["engine_stats_stream", "dynamo_relay_kv_usage"]); + } + + #[test] + fn relay_exact_input_tps_is_sticky_across_idle_windows() { + let mut aggregator = test_aggregator_with_initialization( + StatsCollectorConfig::default(), + ModelStatsInitialization::ConfiguredInputTps { + input_tps: 25.0, + pin: false, + }, + ); + let stats = published_stats( + aggregator.apply_update(StatsAggregatorUpdate::RelayLoad(relay_load(Some(40.0), 1))), + ); + assert_eq!(stats.last_mean_input_tps, 40.0); + + let stats = published_stats( + aggregator.apply_update(StatsAggregatorUpdate::RelayLoad(relay_load(Some(0.0), 2))), + ); + assert_eq!(stats.last_mean_input_tps, 40.0); + } + + #[test] + fn pinned_input_tps_overrides_relay_measurements() { + let mut aggregator = test_aggregator_with_initialization( + StatsCollectorConfig::default(), + ModelStatsInitialization::ConfiguredInputTps { + input_tps: 25.0, + pin: true, + }, + ); + let stats = published_stats( + aggregator.apply_update(StatsAggregatorUpdate::RelayLoad(relay_load(Some(40.0), 1))), + ); + assert_eq!(stats.last_mean_input_tps, 25.0); } #[test] @@ -1458,7 +1502,7 @@ mod tests { .pop() .expect("engine counters should publish the complete owned snapshot") .1; - assert_stats!(stats; output_tps: 10.0, num_running_queries: 1, output_generation_queries: 1, kv_cache_capacity_tokens: 1_000, stats_sources: ["engine_stats_stream", "kv_cache_stats"]); + assert_stats!(stats; output_tps: 10.0, num_running_queries: 1, output_generation_queries: 1, kv_cache_capacity_tokens: 1_000, stats_sources: ["engine_stats_stream", "dynamo_relay_kv_usage"]); } #[test] @@ -1515,41 +1559,6 @@ mod tests { assert_stats!(stats; last_mean_input_tps: 100.0, embedding_item_tps: 2.0, max_embedding_item_tps: 2.0); } - #[tokio::test] - async fn stats_collector_enables_openai_fallback_only_after_control_update() { - let metrics = PylonMetrics::new().expect("metrics should initialize"); - let config = config!(collector; openai_fallback_stats_enabled: false); - let collector = RunningCollector::spawn(config, Some(metrics.clone()), true); - let stats = collector - .observe_until( - trusted_completed_observation("req-fallback-disabled"), - "fallback-disabled observation should publish lifecycle-only stats", - |_| true, - ) - .await; - assert_eq!(stats.output_tps, 0.0); - assert!(!stats.stats_sources.contains(&"chunk_usage".to_string())); - collector - .send_update(StatsAggregatorUpdate::EnableOpenAiFallback) - .await; - wait_for_metric( - &metrics, - r#"pylon_engine_stats_source_transitions_total{from="engine_stats_stream",reason="unsupported",to="openai_fallback"} 1"#, - "collector should process fallback control update before fallback observations are accepted", - ) - .await; - let stats = collector - .observe_until( - trusted_completed_observation("req-fallback-enabled"), - "fallback-enabled observation should publish model stats", - |stats| stats.output_tps == 5.0, - ) - .await; - assert_eq!(stats.output_tps, 5.0); - assert!(stats.stats_sources.contains(&"chunk_usage".to_string())); - collector.handle.shutdown().await; - } - #[tokio::test] async fn stats_collector_keeps_lifecycle_load_when_fallback_stats_disabled() { let config = config!(collector; openai_fallback_stats_enabled: false); @@ -2272,8 +2281,7 @@ mod tests { let mut aggregator = test_aggregator(StatsCollectorConfig::default()); let mut snapshot = kv_cache_stats("canonical-model"); snapshot.aliases.push("model-a".to_string()); - let updates = - aggregator.apply_kv_cache_snapshot(kv_cache_envelope(snapshot), TokioInstant::now()); + let updates = aggregator.apply_kv_cache_snapshot(kv_cache_envelope(snapshot)); let stats = published_stats(updates); assert_stats!(stats; kv_cache_capacity_tokens: 1_000, kv_cache_used_tokens: 400, kv_cache_free_tokens: 600); assert!(stats.kv_cache.is_some()); @@ -2281,31 +2289,12 @@ mod tests { let mut incomplete = kv_cache_stats("canonical-model"); incomplete.aliases.push("model-a".to_string()); incomplete.complete = false; - let updates = - aggregator.apply_kv_cache_snapshot(kv_cache_envelope(incomplete), TokioInstant::now()); + let updates = aggregator.apply_kv_cache_snapshot(kv_cache_envelope(incomplete)); let stats = published_stats(updates); assert_stats!(stats; kv_cache_capacity_tokens: 0, kv_cache_used_tokens: 0, kv_cache_free_tokens: 0); assert!(stats.kv_cache.is_none()); } - #[test] - fn stale_kv_cache_snapshot_expires_without_clearing_request_stats() { - let config = config!(kv_cache_stats_ttl: milliseconds(10)); - let mut aggregator = test_aggregator(config); - aggregator.stream("req", (0, 0), false, Duration::ZERO); - let request_updates = aggregator.stream("req", (10, 4), false, milliseconds(100)); - assert_eq!(published_stats(request_updates).output_tps, 40.0); - - let received_at = TokioInstant::now(); - aggregator - .apply_kv_cache_snapshot(kv_cache_envelope(kv_cache_stats("model-a")), received_at); - let updates = aggregator.sweep_stale(received_at + milliseconds(11)); - let stats = published_stats(updates); - assert!(stats.kv_cache.is_none()); - assert_eq!(stats.kv_cache_capacity_tokens, 0); - assert_eq!(stats.output_tps, 40.0); - } - #[tokio::test] async fn mixed_stream_kv_update_updates_model_metrics() { let metrics = PylonMetrics::new().expect("metrics should initialize"); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs index 584f030ac..d44e57ffa 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs @@ -3,27 +3,24 @@ use std::fmt; use std::str::FromStr; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use std::time::Duration; -use stargate_proto::dynamo_frontend_stats as proto; +use stargate_proto::dynamo_kv_dc_relay as proto; use stargate_runtime::OwnedTask; -use tokio::time::Instant as TokioInstant; use tokio_util::sync::CancellationToken; -use tonic::Code; -use super::collector::{RequestCounterUpdate, StatsAggregatorUpdate, StatsUpdateSource}; -use super::kv_stats::kv_snapshot_from_proto; +use super::collector::StatsAggregatorUpdate; +use super::kv_stats::{kv_snapshot_from_proto, load_snapshot_from_proto}; use super::metrics::PylonMetrics; use crate::PylonRuntimeState; -use crate::generated_request_id::generated_request_generation; const DEFAULT_INITIAL_RECONNECT_BACKOFF: Duration = Duration::from_millis(100); const DEFAULT_MAX_RECONNECT_BACKOFF: Duration = Duration::from_secs(5); +const RELAY_SILENCE_TIMEOUT: Duration = Duration::from_secs(3); #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum EngineStatsStreamMode { - Auto, Required, Off, } @@ -31,7 +28,6 @@ pub enum EngineStatsStreamMode { impl EngineStatsStreamMode { pub fn as_str(self) -> &'static str { match self { - Self::Auto => "auto", Self::Required => "required", Self::Off => "off", } @@ -39,8 +35,8 @@ impl EngineStatsStreamMode { } impl fmt::Display for EngineStatsStreamMode { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.write_str(self.as_str()) + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(self.as_str()) } } @@ -49,7 +45,6 @@ impl FromStr for EngineStatsStreamMode { fn from_str(value: &str) -> Result { match value { - "auto" => Ok(Self::Auto), "required" => Ok(Self::Required), "off" => Ok(Self::Off), _ => Err(ParseEngineStatsStreamModeError), @@ -58,7 +53,7 @@ impl FromStr for EngineStatsStreamMode { } #[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] -#[error("expected one of auto, required, off")] +#[error("expected one of required, off")] pub struct ParseEngineStatsStreamModeError; #[derive(Debug, Clone)] @@ -72,9 +67,9 @@ pub struct EngineStatsStreamConfig { } impl EngineStatsStreamConfig { - pub fn new(upstream_base_url: &str, mode: EngineStatsStreamMode) -> Self { + pub fn new(relay_endpoint: &str, mode: EngineStatsStreamMode) -> Self { Self { - endpoint: upstream_base_url.trim_end_matches('/').to_string(), + endpoint: relay_endpoint.trim_end_matches('/').to_string(), mode, initial_reconnect_backoff: DEFAULT_INITIAL_RECONNECT_BACKOFF, max_reconnect_backoff: DEFAULT_MAX_RECONNECT_BACKOFF, @@ -86,7 +81,7 @@ impl EngineStatsStreamConfig { impl Default for EngineStatsStreamConfig { fn default() -> Self { - Self::new("http://127.0.0.1:8090", EngineStatsStreamMode::Auto) + Self::new("http://127.0.0.1:50051", EngineStatsStreamMode::Required) } } @@ -99,191 +94,301 @@ pub fn start_engine_stats_stream( if config.mode == EngineStatsStreamMode::Off { return None; } - let task = OwnedTask::spawn("engine stats stream", move |stop| { + if let Some(runtime_state) = &config.runtime_state { + runtime_state.require_relay_load(true); + runtime_state.mark_relay_load_unavailable(); + } + let task = OwnedTask::spawn("Dynamo Relay stats streams", move |stop| { run_engine_stats_stream(config, stats_update_tx, stop) }); Some(EngineStatsStreamHandle { task }) } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct RelayIdentity { + drt_instance_id: u64, + relay_incarnation: u64, +} + +#[derive(Default)] +struct RelayEpoch { + load: Option, +} + async fn run_engine_stats_stream( config: EngineStatsStreamConfig, stats_update_tx: flume::Sender, stop: CancellationToken, +) { + let epoch = Arc::new(Mutex::new(RelayEpoch::default())); + tokio::join!( + run_load_stream( + config.clone(), + stats_update_tx.clone(), + epoch.clone(), + stop.clone(), + ), + run_kv_stream(config, stats_update_tx, epoch, stop), + ); +} + +async fn run_load_stream( + config: EngineStatsStreamConfig, + updates: flume::Sender, + epoch: Arc>, + stop: CancellationToken, ) { let mut backoff = config.initial_reconnect_backoff; - let mut valid_event_seen = false; loop { if stop.is_cancelled() { return; } - let retry_reason = - match read_stream_once(&config, &stats_update_tx, &stop, &mut valid_event_seen).await { - StreamReadOutcome::Stopped => return, - StreamReadOutcome::Unsupported - if config.mode == EngineStatsStreamMode::Auto && !valid_event_seen => - { + mark_load_unavailable(&config, &updates, &epoch, &stop).await; + let connect = proto::kv_dc_relay_client::KvDcRelayClient::connect(config.endpoint.clone()); + let mut client = match stop.run_until_cancelled(connect).await { + None => return, + Some(Ok(client)) => client, + Some(Err(error)) => { + tracing::warn!(endpoint = config.endpoint, %error, "Dynamo Relay load connect failed"); + reconnect_delay(&config, &stop, &mut backoff, "load_connect").await; + continue; + } + }; + let response = match stop.run_until_cancelled(client.watch_load(())).await { + None => return, + Some(Ok(response)) => response, + Some(Err(status)) => { + tracing::warn!(endpoint = config.endpoint, %status, "Dynamo Relay load stream failed"); + reconnect_delay(&config, &stop, &mut backoff, "load_watch").await; + continue; + } + }; + + observe_connected(&config, true); + backoff = config.initial_reconnect_backoff; + let mut stream = response.into_inner(); + loop { + let message = tokio::select! { + _ = stop.cancelled() => return, + message = tokio::time::timeout(RELAY_SILENCE_TIMEOUT, stream.message()) => message, + }; + let snapshot = match message { + Ok(Ok(Some(snapshot))) => snapshot, + Ok(Ok(None)) => break, + Ok(Err(status)) => { + tracing::warn!(endpoint = config.endpoint, %status, "Dynamo Relay load stream disconnected"); + break; + } + Err(_) => { tracing::warn!( endpoint = config.endpoint, - "frontend stats gRPC service unsupported; using OpenAI fallback observation" + "Dynamo Relay load stream became stale" ); - let _ = send_stats_update( - &stats_update_tx, - StatsAggregatorUpdate::EnableOpenAiFallback, + break; + } + }; + let identity = match metadata_identity(snapshot.metadata.as_ref()) { + Ok(identity) => identity, + Err(error) => { + invalid_event(&config, "load_metadata", error); + mark_load_unavailable(&config, &updates, &epoch, &stop).await; + continue; + } + }; + let identity_changed = { + let mut epoch = epoch.lock().expect("Relay epoch mutex poisoned"); + let changed = epoch.load.is_some_and(|current| current != identity); + epoch.load = Some(identity); + changed + }; + if identity_changed { + send_update( + &updates, + StatsAggregatorUpdate::KvCache(Default::default()), + &stop, + ) + .await; + } + match load_snapshot_from_proto(snapshot) { + Ok(translation) => { + if let Some(runtime_state) = &config.runtime_state { + runtime_state.replace_relay_models(translation.relay_models); + } + if !send_update( + &updates, + StatsAggregatorUpdate::RelayLoad(translation.stats), &stop, ) - .await; - return; + .await + { + return; + } + observe_event(&config, "load_snapshot"); } - StreamReadOutcome::Unsupported => "unsupported", - StreamReadOutcome::Retry(reason) => reason, - }; - - if let Some(metrics) = &config.metrics { - metrics.observe_engine_stats_reconnect(retry_reason); - } - if valid_event_seen { - backoff = config.initial_reconnect_backoff; - } - if stop - .run_until_cancelled(tokio::time::sleep(backoff)) - .await - .is_none() - { - return; + Err(error) => { + invalid_event(&config, "load_snapshot", error); + mark_load_unavailable(&config, &updates, &epoch, &stop).await; + } + } } - backoff = (backoff * 2).min(config.max_reconnect_backoff); + observe_connected(&config, false); + mark_load_unavailable(&config, &updates, &epoch, &stop).await; + reconnect_delay(&config, &stop, &mut backoff, "load_eof").await; } } -#[derive(Debug)] -enum StreamReadOutcome { - Stopped, - Unsupported, - Retry(&'static str), -} - -async fn read_stream_once( - config: &EngineStatsStreamConfig, - stats_update_tx: &flume::Sender, - stop: &CancellationToken, - valid_event_seen: &mut bool, -) -> StreamReadOutcome { - let connect = - proto::frontend_stats_client::FrontendStatsClient::connect(config.endpoint.clone()); - let mut client = match stop.run_until_cancelled(connect).await { - None => return StreamReadOutcome::Stopped, - Some(Ok(client)) => client, - Some(Err(error)) => { - tracing::warn!(endpoint = config.endpoint, %error, "frontend stats gRPC connect failed"); - return StreamReadOutcome::Retry("connect_error"); +async fn run_kv_stream( + config: EngineStatsStreamConfig, + updates: flume::Sender, + epoch: Arc>, + stop: CancellationToken, +) { + let mut backoff = config.initial_reconnect_backoff; + loop { + if stop.is_cancelled() { + return; } - }; - - let watch = client.watch_stats(proto::WatchStatsRequest {}); - let response = match stop.run_until_cancelled(watch).await { - None => return StreamReadOutcome::Stopped, - Some(Ok(response)) => response, - Some(Err(status)) => return status_outcome(config, status), - }; - - observe_connected(config, true); - let mut stream = response.into_inner(); - let outcome = loop { - let message = match stop.run_until_cancelled(stream.message()).await { - None => break StreamReadOutcome::Stopped, - Some(message) => message, + send_update( + &updates, + StatsAggregatorUpdate::KvCache(Default::default()), + &stop, + ) + .await; + let connect = proto::kv_dc_relay_client::KvDcRelayClient::connect(config.endpoint.clone()); + let mut client = match stop.run_until_cancelled(connect).await { + None => return, + Some(Ok(client)) => client, + Some(Err(error)) => { + tracing::warn!(endpoint = config.endpoint, %error, "Dynamo Relay KV-usage connect failed"); + reconnect_delay(&config, &stop, &mut backoff, "kv_connect").await; + continue; + } }; - let update = match message { - Ok(Some(update)) => update, - Ok(None) => break StreamReadOutcome::Retry("eof"), - Err(status) => break status_outcome(config, status), + let response = match stop.run_until_cancelled(client.watch_kv_usage(())).await { + None => return, + Some(Ok(response)) => response, + Some(Err(status)) => { + tracing::warn!(endpoint = config.endpoint, %status, "Dynamo Relay KV-usage stream failed"); + reconnect_delay(&config, &stop, &mut backoff, "kv_watch").await; + continue; + } }; - let (event_type, update) = match translate_update(config, update, TokioInstant::now()) { - Ok(update) => update, - Err(error) => { - tracing::warn!(endpoint = config.endpoint, %error, "invalid frontend stats update"); - if let Some(metrics) = &config.metrics { - metrics.observe_engine_stats_invalid_event("protobuf"); + backoff = config.initial_reconnect_backoff; + let mut stream = response.into_inner(); + loop { + let message = tokio::select! { + _ = stop.cancelled() => return, + message = tokio::time::timeout(RELAY_SILENCE_TIMEOUT, stream.message()) => message, + }; + let snapshot = match message { + Ok(Ok(Some(snapshot))) => snapshot, + Ok(Ok(None)) => break, + Ok(Err(status)) => { + tracing::warn!(endpoint = config.endpoint, %status, "Dynamo Relay KV-usage stream disconnected"); + break; + } + Err(_) => { + tracing::warn!( + endpoint = config.endpoint, + "Dynamo Relay KV-usage stream became stale" + ); + break; } + }; + let identity = match metadata_identity(snapshot.metadata.as_ref()) { + Ok(identity) => identity, + Err(error) => { + invalid_event(&config, "kv_metadata", error); + clear_kv(&updates, &stop).await; + continue; + } + }; + if epoch.lock().expect("Relay epoch mutex poisoned").load != Some(identity) { + clear_kv(&updates, &stop).await; continue; } - }; - *valid_event_seen = true; - if let Some(metrics) = &config.metrics { - metrics.observe_engine_stats_stream_event(event_type); - } - if !send_stats_update(stats_update_tx, update, stop).await { - break StreamReadOutcome::Stopped; + match kv_snapshot_from_proto(snapshot) { + Ok(snapshot) => { + if !send_update(&updates, StatsAggregatorUpdate::KvCache(snapshot), &stop).await + { + return; + } + observe_event(&config, "kv_usage_snapshot"); + } + Err(error) => { + invalid_event(&config, "kv_usage_snapshot", error); + clear_kv(&updates, &stop).await; + } + } } - }; - observe_connected(config, false); - outcome + clear_kv(&updates, &stop).await; + reconnect_delay(&config, &stop, &mut backoff, "kv_eof").await; + } } -fn status_outcome(config: &EngineStatsStreamConfig, status: tonic::Status) -> StreamReadOutcome { - if matches!(status.code(), Code::Unimplemented | Code::NotFound) { - tracing::warn!(endpoint = config.endpoint, %status, "frontend stats gRPC service is unsupported"); - StreamReadOutcome::Unsupported - } else { - tracing::warn!(endpoint = config.endpoint, %status, "frontend stats gRPC stream disconnected"); - StreamReadOutcome::Retry("grpc_status") +async fn mark_load_unavailable( + config: &EngineStatsStreamConfig, + updates: &flume::Sender, + epoch: &Mutex, + stop: &CancellationToken, +) { + epoch.lock().expect("Relay epoch mutex poisoned").load = None; + if let Some(runtime_state) = &config.runtime_state { + runtime_state.mark_relay_load_unavailable(); } + send_update( + updates, + StatsAggregatorUpdate::RelayLoad(Default::default()), + stop, + ) + .await; + clear_kv(updates, stop).await; } -fn translate_update( +async fn clear_kv(updates: &flume::Sender, stop: &CancellationToken) { + send_update( + updates, + StatsAggregatorUpdate::KvCache(Default::default()), + stop, + ) + .await; +} + +fn metadata_identity( + metadata: Option<&proto::RelayMessageMetadata>, +) -> anyhow::Result { + let metadata = metadata.ok_or_else(|| anyhow::anyhow!("Relay metadata is missing"))?; + anyhow::ensure!(metadata.drt_instance_id != 0, "DRT instance ID is zero"); + anyhow::ensure!(metadata.relay_incarnation != 0, "Relay incarnation is zero"); + Ok(RelayIdentity { + drt_instance_id: metadata.drt_instance_id, + relay_incarnation: metadata.relay_incarnation, + }) +} + +async fn reconnect_delay( config: &EngineStatsStreamConfig, - update: proto::StatsUpdate, - observed_at: TokioInstant, -) -> anyhow::Result<(&'static str, StatsAggregatorUpdate)> { - match update.update { - Some(proto::stats_update::Update::RequestStats(request)) => { - let request_id = request.request_id.trim(); - let model_id = request.model.trim(); - anyhow::ensure!(!request_id.is_empty(), "request ID is empty"); - anyhow::ensure!(!model_id.is_empty(), "model ID is empty"); - anyhow::ensure!( - request.tokens_processed.is_some() - || request.tokens_generated.is_some() - || request.finished, - "request update has no counters" - ); - let generation = generated_request_generation(request_id, model_id).or_else(|| { - config - .runtime_state - .as_ref() - .and_then(|state| state.request_generation(request_id)) - }); - Ok(( - "stats", - StatsAggregatorUpdate::RequestCounters(RequestCounterUpdate { - source: StatsUpdateSource::EngineStatsStream, - request_id: request_id.to_string(), - model_id: model_id.to_string(), - generation, - tokens_processed: request.tokens_processed, - tokens_generated: request.tokens_generated, - finished: request.finished, - observed_at, - }), - )) - } - Some(proto::stats_update::Update::KvStats(snapshot)) => Ok(( - "kv_stats_snapshot", - StatsAggregatorUpdate::KvCache(kv_snapshot_from_proto(snapshot)?), - )), - None => anyhow::bail!("stats update is missing its payload"), + stop: &CancellationToken, + backoff: &mut Duration, + reason: &'static str, +) { + if let Some(metrics) = &config.metrics { + metrics.observe_engine_stats_reconnect(reason); } + let delay = *backoff; + *backoff = (*backoff * 2).min(config.max_reconnect_backoff); + let _ = stop.run_until_cancelled(tokio::time::sleep(delay)).await; } -async fn send_stats_update( - stats_update_tx: &flume::Sender, +async fn send_update( + sender: &flume::Sender, update: StatsAggregatorUpdate, stop: &CancellationToken, ) -> bool { - match stats_update_tx.try_send(update) { + match sender.try_send(update) { Ok(()) => true, Err(flume::TrySendError::Full(update)) => stop - .run_until_cancelled(stats_update_tx.send_async(update)) + .run_until_cancelled(sender.send_async(update)) .await .is_some_and(|result| result.is_ok()), Err(flume::TrySendError::Disconnected(_)) => false, @@ -296,204 +401,34 @@ fn observe_connected(config: &EngineStatsStreamConfig, connected: bool) { } } +fn observe_event(config: &EngineStatsStreamConfig, event: &'static str) { + if let Some(metrics) = &config.metrics { + metrics.observe_engine_stats_stream_event(event); + } +} + +fn invalid_event(config: &EngineStatsStreamConfig, kind: &'static str, error: anyhow::Error) { + tracing::warn!(endpoint = config.endpoint, %error, "invalid Dynamo Relay stats update"); + if let Some(metrics) = &config.metrics { + metrics.observe_engine_stats_invalid_event(kind); + } +} + #[cfg(test)] mod tests { use super::*; - use std::pin::Pin; - use std::sync::atomic::{AtomicUsize, Ordering}; - - use futures::Stream; - use tokio::net::TcpListener; - - fn request_update(request: proto::RequestStats) -> proto::StatsUpdate { - proto::StatsUpdate { - update: Some(proto::stats_update::Update::RequestStats(request)), - } - } #[test] fn parses_stream_modes() { - assert_eq!("auto".parse(), Ok(EngineStatsStreamMode::Auto)); assert_eq!("required".parse(), Ok(EngineStatsStreamMode::Required)); assert_eq!("off".parse(), Ok(EngineStatsStreamMode::Off)); + assert!("auto".parse::().is_err()); assert!("other".parse::().is_err()); } #[test] - fn translates_request_counters() { - let update = request_update(proto::RequestStats { - request_id: "request-a".to_string(), - model: "model-a".to_string(), - tokens_processed: Some(12), - tokens_generated: Some(3), - finished: false, - }); - let (_, update) = translate_update( - &EngineStatsStreamConfig::default(), - update, - TokioInstant::now(), - ) - .unwrap(); - let StatsAggregatorUpdate::RequestCounters(update) = update else { - panic!("expected request counters"); - }; - assert_eq!(update.request_id, "request-a"); - assert_eq!(update.model_id, "model-a"); - assert_eq!(update.tokens_processed, Some(12)); - assert_eq!(update.tokens_generated, Some(3)); - } - - #[test] - fn rejects_empty_request_updates() { - let update = request_update(proto::RequestStats { - request_id: "request-a".to_string(), - model: "model-a".to_string(), - tokens_processed: None, - tokens_generated: None, - finished: false, - }); - assert!( - translate_update( - &EngineStatsStreamConfig::default(), - update, - TokioInstant::now() - ) - .is_err() - ); - } - - #[test] - fn translates_kv_snapshots_on_the_same_stream() { - let update = proto::StatsUpdate { - update: Some(proto::stats_update::Update::KvStats( - proto::KvStatsSnapshot { - snapshot_id: 1, - observed_at_unix_ms: 10, - models: Vec::new(), - }, - )), - }; - let (_, update) = translate_update( - &EngineStatsStreamConfig::default(), - update, - TokioInstant::now(), - ) - .unwrap(); - assert!(matches!(update, StatsAggregatorUpdate::KvCache(_))); - } - - #[derive(Clone)] - struct TestFrontendStats { - calls: Arc, - updates: Arc>, - } - - #[tonic::async_trait] - impl proto::frontend_stats_server::FrontendStats for TestFrontendStats { - type WatchStatsStream = - Pin> + Send>>; - type WatchKvPlacementsStream = - Pin> + Send>>; - - async fn watch_stats( - &self, - _request: tonic::Request, - ) -> Result, tonic::Status> { - self.calls.fetch_add(1, Ordering::Relaxed); - let stream = futures::stream::iter(self.updates.as_ref().clone().into_iter().map(Ok)); - Ok(tonic::Response::new(Box::pin(stream))) - } - - async fn watch_kv_placements( - &self, - _request: tonic::Request, - ) -> Result, tonic::Status> { - Err(tonic::Status::unimplemented("not used")) - } - } - - async fn spawn_test_server(router: axum::Router) -> (String, tokio::task::JoinHandle<()>) { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let endpoint = format!("http://{}", listener.local_addr().unwrap()); - let server = tokio::spawn(async move { - axum::serve(listener, router).await.unwrap(); - }); - (endpoint, server) - } - - #[tokio::test] - async fn grpc_stream_delivers_both_updates_and_reconnects_after_eof() { - let calls = Arc::new(AtomicUsize::new(0)); - let updates = Arc::new(vec![ - request_update(proto::RequestStats { - request_id: "request-a".to_string(), - model: "model-a".to_string(), - tokens_processed: Some(10), - tokens_generated: None, - finished: false, - }), - proto::StatsUpdate { - update: Some(proto::stats_update::Update::KvStats( - proto::KvStatsSnapshot { - snapshot_id: 1, - observed_at_unix_ms: 10, - models: Vec::new(), - }, - )), - }, - ]); - let service = proto::frontend_stats_server::FrontendStatsServer::new(TestFrontendStats { - calls: calls.clone(), - updates, - }); - let router = tonic::service::Routes::new(service).into_axum_router(); - let (endpoint, server) = spawn_test_server(router).await; - let config = EngineStatsStreamConfig { - initial_reconnect_backoff: Duration::from_millis(1), - max_reconnect_backoff: Duration::from_millis(1), - ..EngineStatsStreamConfig::new(&endpoint, EngineStatsStreamMode::Required) - }; - let (tx, rx) = flume::bounded(8); - let stream = start_engine_stats_stream(config, tx).unwrap(); - - let first = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) - .await - .unwrap() - .unwrap(); - let second = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) - .await - .unwrap() - .unwrap(); - let third = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) - .await - .unwrap() - .unwrap(); - assert!(matches!(first, StatsAggregatorUpdate::RequestCounters(_))); - assert!(matches!(second, StatsAggregatorUpdate::KvCache(_))); - assert!(matches!(third, StatsAggregatorUpdate::RequestCounters(_))); - assert!(calls.load(Ordering::Relaxed) >= 2); - - stream.shutdown().await; - server.abort(); - } - - #[tokio::test] - async fn auto_mode_falls_back_when_grpc_service_is_absent() { - let (endpoint, server) = spawn_test_server(axum::Router::new()).await; - let config = EngineStatsStreamConfig::new(&endpoint, EngineStatsStreamMode::Auto); - let (tx, rx) = flume::bounded(1); - let stream = start_engine_stats_stream(config, tx).unwrap(); - - let update = tokio::time::timeout(Duration::from_secs(1), rx.recv_async()) - .await - .expect("auto mode should resolve unsupported gRPC") - .unwrap(); - assert!(matches!( - update, - StatsAggregatorUpdate::EnableOpenAiFallback - )); - - stream.shutdown().await; - server.abort(); + fn rejects_missing_or_zero_relay_incarnation() { + assert!(metadata_identity(None).is_err()); + assert!(metadata_identity(Some(&proto::RelayMessageMetadata::default())).is_err()); } } diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs index 9e0d734ba..28eb41fcc 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs @@ -1,113 +1,591 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -//! Translation of Dynamo frontend KV snapshots into Pylon's aggregate model state. +//! Translation from Dynamo's Relay-owned contract into Pylon model state. -use std::collections::HashSet; +use std::collections::{BTreeMap, HashMap, HashSet}; -use stargate_proto::dynamo_frontend_stats as proto; +use stargate_proto::dynamo_kv_dc_relay as proto; -use super::aggregator::{KvCacheStatsEnvelope, KvCacheStatsSnapshot}; +use super::aggregator::{ + KvCacheStatsEnvelope, KvCacheStatsSnapshot, RelayLoadStatsEnvelope, RelayLoadStatsSnapshot, +}; + +#[derive(Clone, Debug, Hash, PartialEq, Eq)] +struct PoolKey { + cache_semantics: [u8; 16], + cache_source: i32, + routing_scope: [u8; 16], + routing_source: i32, + dc_id: u64, +} + +#[derive(Clone)] +struct UsagePool { + role: proto::WorkerRole, + identities: Vec, + capacity_tokens: Option, + used_tokens: Option, + source_observed_at_unix_ms: u64, + complete: bool, +} + +#[derive(Clone)] +struct LoadPool { + role: proto::WorkerRole, + live_workers: Option, + max_concurrency: Option, + complete: bool, +} + +pub(super) struct RelayLoadTranslation { + pub(super) stats: RelayLoadStatsEnvelope, + pub(super) relay_models: BTreeMap, +} pub(super) fn kv_snapshot_from_proto( - snapshot: proto::KvStatsSnapshot, + snapshot: proto::KvUsageSnapshot, ) -> anyhow::Result { - let mut identities = HashSet::new(); - let models = snapshot - .models - .into_iter() - .map(|model| -> anyhow::Result<_> { - anyhow::ensure!(!model.model.trim().is_empty(), "KV stats model is empty"); - anyhow::ensure!( - identities.insert(model.model.clone()), - "duplicate KV stats identity {}", - model.model - ); - for alias in &model.aliases { - anyhow::ensure!(!alias.trim().is_empty(), "KV stats alias is empty"); - anyhow::ensure!( - identities.insert(alias.clone()), - "duplicate KV stats identity {alias}" - ); + snapshot + .metadata + .as_ref() + .ok_or_else(|| anyhow::anyhow!("KV usage snapshot metadata is missing"))?; + let mut pools = Vec::with_capacity(snapshot.pools.len()); + let mut pool_ids = HashSet::new(); + let mut identity_owner = HashMap::::new(); + for pool in snapshot.pools { + let key = pool_key(pool.pool)?; + anyhow::ensure!(pool_ids.insert(key.clone()), "duplicate KV usage pool"); + let role = worker_role(pool.role)?; + let identities = registration_identities(&pool.models, &mut identity_owner)?; + let complete = data_complete(pool.status) + && pool.expected_ranks > 0 + && pool.observed_ranks == pool.expected_ranks + && pool.block_size_tokens > 0; + let (capacity_tokens, used_tokens) = match (pool.capacity_blocks, pool.used_blocks) { + (Some(capacity), Some(used)) if used <= capacity => ( + capacity.checked_mul(u64::from(pool.block_size_tokens)), + used.checked_mul(u64::from(pool.block_size_tokens)), + ), + _ => (None, None), + }; + pools.push(UsagePool { + role, + identities, + capacity_tokens, + used_tokens, + source_observed_at_unix_ms: pool.source_observed_at_unix_ms, + complete: complete + && pool.source_observed_at_unix_ms > 0 + && capacity_tokens.is_some() + && used_tokens.is_some(), + }); + } + + let mut by_identity = BTreeMap::>::new(); + for (index, pool) in pools.iter().enumerate() { + for identity in &pool.identities { + let indexes = by_identity.entry(identity.clone()).or_default(); + if !indexes.contains(&index) { + indexes.push(index); } + } + } - let complete = snapshot.observed_at_unix_ms > 0 - && model.routing_cache.as_ref().is_some_and(|routing| { - routing.capacity_tokens > 0 - && routing.used_tokens.checked_add(routing.free_tokens) - == Some(routing.capacity_tokens) - }); - let (capacity, used, free) = model - .routing_cache - .map(|routing| { - ( - routing.capacity_tokens, - routing.used_tokens, - routing.free_tokens, - ) + let models = by_identity + .into_iter() + .map(|(model, indexes)| { + let role = if indexes + .iter() + .any(|index| pools[*index].role == proto::WorkerRole::Aggregated) + { + proto::WorkerRole::Aggregated + } else { + proto::WorkerRole::Decode + }; + let selected = indexes + .into_iter() + .filter(|index| pools[*index].role == role) + .collect::>(); + let complete = + !selected.is_empty() && selected.iter().all(|index| pools[*index].complete); + let totals = complete + .then(|| { + selected + .iter() + .try_fold((0_u64, 0_u64), |(capacity, used), index| { + Some(( + capacity.checked_add(pools[*index].capacity_tokens?)?, + used.checked_add(pools[*index].used_tokens?)?, + )) + }) }) + .flatten(); + let (capacity, used, free) = totals + .and_then(|(capacity, used)| Some((capacity, used, capacity.checked_sub(used)?))) .unwrap_or_default(); - Ok(KvCacheStatsSnapshot { - model: model.model, - aliases: model.aliases, + KvCacheStatsSnapshot { + model, + aliases: Vec::new(), kv_cache_capacity_tokens: capacity, kv_cache_used_tokens: used, kv_cache_free_tokens: free, - source_observed_at_unix_ms: snapshot.observed_at_unix_ms, - complete, - }) + source_observed_at_unix_ms: selected + .iter() + .map(|index| pools[*index].source_observed_at_unix_ms) + .filter(|timestamp| *timestamp > 0) + .min() + .unwrap_or_default(), + complete: complete && totals.is_some(), + } }) - .collect::>>()?; + .collect(); Ok(KvCacheStatsEnvelope { models }) } +pub(super) fn load_snapshot_from_proto( + snapshot: proto::LoadSnapshot, +) -> anyhow::Result { + anyhow::ensure!(snapshot.window_ms > 0, "load snapshot window is zero"); + snapshot + .metadata + .as_ref() + .ok_or_else(|| anyhow::anyhow!("load snapshot metadata is missing"))?; + let window_seconds = f64::from(snapshot.window_ms) / 1_000.0; + let mut pools = HashMap::::new(); + for pool in snapshot.pools { + let key = pool_key(pool.pool)?; + let role = worker_role(pool.role)?; + anyhow::ensure!( + pools + .insert( + key, + LoadPool { + role, + live_workers: pool.live_workers, + max_concurrency: pool.max_concurrency, + complete: data_complete(pool.scheduler_status), + }, + ) + .is_none(), + "duplicate load pool" + ); + } + + let mut identity_owner = HashMap::::new(); + let mut relay_models = BTreeMap::new(); + let mut models = Vec::new(); + for model in snapshot.models { + let registration = model + .model + .as_ref() + .ok_or_else(|| anyhow::anyhow!("load model registration is missing"))?; + let identities = + registration_identities(std::slice::from_ref(registration), &mut identity_owner)?; + let serving_pools = model + .serving_pools + .into_iter() + .map(|pool| pool_key(Some(pool))) + .collect::>>()?; + anyhow::ensure!( + serving_pools.iter().all(|pool| pools.contains_key(pool)), + "load model references an unknown serving pool" + ); + let selected_role = if serving_pools.iter().any(|pool| { + pools + .get(pool) + .is_some_and(|pool| pool.role == proto::WorkerRole::Aggregated) + }) { + proto::WorkerRole::Aggregated + } else { + proto::WorkerRole::Decode + }; + let selected = serving_pools + .iter() + .filter_map(|key| pools.get(key)) + .filter(|pool| pool.role == selected_role) + .collect::>(); + let scheduler_live = selected + .iter() + .any(|pool| pool.complete && pool.live_workers.is_some_and(|workers| workers > 0)); + let max_engine_concurrency = (!selected.is_empty() + && selected + .iter() + .all(|pool| pool.complete && pool.max_concurrency.is_some())) + .then(|| { + selected.iter().try_fold(0_u64, |total, pool| { + total.checked_add(pool.max_concurrency?) + }) + }) + .flatten(); + let required = ( + model.ready_frontends, + model.pending_first_output_requests, + model.input_processing_requests, + model.output_generation_requests, + ); + let complete = data_complete(model.status) + && model.expected_frontends > 0 + && model.observed_frontends == model.expected_frontends + && model.source_observed_at_unix_ms > 0 + && matches!(required, (Some(_), Some(_), Some(_), Some(_))); + let (ready_frontends, queue_size, input_processing_queries, output_generation_queries) = + required; + let num_running_queries = input_processing_queries + .zip(output_generation_queries) + .and_then(|(input, output)| input.checked_add(output)); + let complete = complete && num_running_queries.is_some(); + let active = complete + && !serving_pools.is_empty() + && ready_frontends.is_some_and(|ready| ready > 0) + && scheduler_live; + for identity in &identities { + relay_models.insert(identity.clone(), active); + models.push(RelayLoadStatsSnapshot { + model: identity.clone(), + input_tps: if complete { + model + .input_tokens + .map(|tokens| tokens as f64 / window_seconds) + } else { + None + }, + output_tps: if complete { + model.output_tokens as f64 / window_seconds + } else { + 0.0 + }, + queue_size: queue_size.unwrap_or_default(), + queued_input_size: complete + .then_some(model.pending_first_output_input_tokens) + .flatten(), + num_running_queries: num_running_queries.unwrap_or_default(), + max_engine_concurrency, + total_query_input_size: complete.then_some(model.live_input_tokens).flatten(), + input_processing_queries: input_processing_queries.unwrap_or_default(), + output_generation_queries: output_generation_queries.unwrap_or_default(), + source_observed_at_unix_ms: model.source_observed_at_unix_ms, + complete, + }); + } + } + Ok(RelayLoadTranslation { + stats: RelayLoadStatsEnvelope { models }, + relay_models, + }) +} + +fn pool_key(pool: Option) -> anyhow::Result { + let pool = pool.ok_or_else(|| anyhow::anyhow!("pool identity is missing"))?; + let cache_semantics: [u8; 16] = pool + .cache_semantics_digest + .try_into() + .map_err(|_| anyhow::anyhow!("cache-semantics digest must contain 16 bytes"))?; + let routing_scope: [u8; 16] = pool + .routing_scope_digest + .try_into() + .map_err(|_| anyhow::anyhow!("routing-scope digest must contain 16 bytes"))?; + identity_source(pool.cache_semantics_source)?; + identity_source(pool.routing_scope_source)?; + Ok(PoolKey { + cache_semantics, + cache_source: pool.cache_semantics_source, + routing_scope, + routing_source: pool.routing_scope_source, + dc_id: pool.dc_id, + }) +} + +fn identity_source(value: i32) -> anyhow::Result { + let source = proto::IdentitySource::try_from(value) + .map_err(|_| anyhow::anyhow!("invalid pool identity source"))?; + anyhow::ensure!( + matches!( + source, + proto::IdentitySource::DefaultDerived | proto::IdentitySource::Explicit + ), + "unspecified pool identity source" + ); + Ok(source) +} + +fn worker_role(value: i32) -> anyhow::Result { + match proto::WorkerRole::try_from(value) { + Ok( + role @ (proto::WorkerRole::Aggregated + | proto::WorkerRole::Prefill + | proto::WorkerRole::Decode + | proto::WorkerRole::Encode), + ) => Ok(role), + _ => anyhow::bail!("invalid or unspecified worker role"), + } +} + +fn data_complete(value: i32) -> bool { + proto::DataStatus::try_from(value).ok() == Some(proto::DataStatus::Complete) +} + +fn registration_identities( + registrations: &[proto::ModelRegistration], + owners: &mut HashMap, +) -> anyhow::Result> { + let mut identities = Vec::new(); + for registration in registrations { + let model = registration.model.trim(); + anyhow::ensure!(!model.is_empty(), "model registration is empty"); + anyhow::ensure!( + !registration.base_model.trim().is_empty(), + "base model registration is empty" + ); + let mut registration_identities = Vec::with_capacity(registration.aliases.len() + 1); + registration_identities.push(model.to_string()); + for alias in ®istration.aliases { + anyhow::ensure!(!alias.trim().is_empty(), "model alias is empty"); + if alias != model && !registration_identities.contains(alias) { + registration_identities.push(alias.clone()); + } + } + for identity in registration_identities { + if let Some(owner) = owners.insert(identity.clone(), model.to_string()) { + anyhow::ensure!( + owner == model, + "model identity {identity} is owned by both {owner} and {model}" + ); + } + if !identities.contains(&identity) { + identities.push(identity); + } + } + } + Ok(identities) +} + #[cfg(test)] mod tests { use super::*; - fn snapshot(capacity: u64, used: u64, free: u64) -> proto::KvStatsSnapshot { - proto::KvStatsSnapshot { - snapshot_id: 1, - observed_at_unix_ms: 10, - models: vec![proto::ModelKvStats { - model: "model-a".to_string(), - aliases: vec!["alias-a".to_string()], - routing_cache: Some(proto::RoutingCacheStats { + fn metadata() -> proto::RelayMessageMetadata { + proto::RelayMessageMetadata { + drt_instance_id: 1, + relay_incarnation: 2, + observed_at_unix_ms: 99, + } + } + + fn pool(seed: u8) -> proto::PoolIdentity { + proto::PoolIdentity { + cache_semantics_digest: vec![seed; 16], + cache_semantics_source: proto::IdentitySource::Explicit as i32, + routing_scope_digest: vec![seed.wrapping_add(1); 16], + routing_scope_source: proto::IdentitySource::DefaultDerived as i32, + dc_id: u64::from(seed), + } + } + + fn registration(model: &str) -> proto::ModelRegistration { + proto::ModelRegistration { + model: model.to_string(), + base_model: model.to_string(), + adapter: None, + aliases: vec![format!("{model}-alias")], + } + } + + fn load_pool(identity: proto::PoolIdentity) -> proto::PoolLoad { + proto::PoolLoad { + pool: Some(identity), + role: proto::WorkerRole::Aggregated as i32, + live_workers: Some(2), + active_prefill_tokens: Some(5), + active_decode_blocks: Some(6), + max_concurrency: Some(8), + scheduler_status: proto::DataStatus::Complete as i32, + scheduler_observed_at_unix_ms: 4, + } + } + + fn model_load(model: &str, serving_pools: Vec) -> proto::ModelLoad { + proto::ModelLoad { + model: Some(registration(model)), + ready_frontends: Some(1), + pending_first_output_requests: Some(2), + pending_first_output_input_tokens: Some(17), + live_input_tokens: Some(31), + input_processing_requests: Some(1), + output_generation_requests: Some(2), + serving_pools, + requests_started: 4, + requests_completed: 3, + requests_failed: 0, + requests_cancelled: 0, + input_tokens: Some(40), + output_tokens: 20, + status: proto::DataStatus::Complete as i32, + expected_frontends: 1, + observed_frontends: 1, + source_observed_at_unix_ms: 5, + } + } + + #[test] + fn kv_usage_prefers_aggregated_pools_and_scales_blocks_to_tokens() { + let aggregated = pool(1); + let decode = pool(2); + let snapshot = proto::KvUsageSnapshot { + metadata: Some(metadata()), + pools: vec![ + proto::PoolKvUsage { + pool: Some(aggregated), + models: vec![registration("model-a")], role: proto::WorkerRole::Aggregated as i32, - capacity_tokens: capacity, - used_tokens: used, - free_tokens: free, - }), - pools: Vec::new(), - }], + block_size_tokens: 16, + expected_ranks: 2, + observed_ranks: 2, + capacity_blocks: Some(100), + used_blocks: Some(40), + status: proto::DataStatus::Complete as i32, + source_observed_at_unix_ms: 7, + }, + proto::PoolKvUsage { + pool: Some(decode), + models: vec![registration("model-a")], + role: proto::WorkerRole::Decode as i32, + block_size_tokens: 16, + expected_ranks: 1, + observed_ranks: 1, + capacity_blocks: Some(1_000), + used_blocks: Some(900), + status: proto::DataStatus::Complete as i32, + source_observed_at_unix_ms: 8, + }, + ], + }; + + let translated = kv_snapshot_from_proto(snapshot).unwrap(); + assert_eq!(translated.models.len(), 2); + for model in translated.models { + assert!(matches!(model.model.as_str(), "model-a" | "model-a-alias")); + assert_eq!(model.kv_cache_capacity_tokens, 1_600); + assert_eq!(model.kv_cache_used_tokens, 640); + assert_eq!(model.kv_cache_free_tokens, 960); + assert_eq!(model.source_observed_at_unix_ms, 7); + assert!(model.complete); } } #[test] - fn converts_complete_routing_cache_stats() { - let envelope = kv_snapshot_from_proto(snapshot(100, 40, 60)).unwrap(); - let model = &envelope.models[0]; - assert!(model.complete); - assert_eq!(model.kv_cache_capacity_tokens, 100); - assert_eq!(model.kv_cache_used_tokens, 40); - assert_eq!(model.kv_cache_free_tokens, 60); + fn complete_load_activates_model_and_alias() { + let identity = pool(1); + let snapshot = proto::LoadSnapshot { + metadata: Some(metadata()), + window_ms: 1_000, + pools: vec![load_pool(identity.clone())], + models: vec![model_load("model-a", vec![identity])], + }; + + let translated = load_snapshot_from_proto(snapshot).unwrap(); + assert_eq!(translated.relay_models.get("model-a"), Some(&true)); + assert_eq!(translated.relay_models.get("model-a-alias"), Some(&true)); + assert_eq!(translated.stats.models.len(), 2); + for model in translated.stats.models { + assert_eq!(model.input_tps, Some(40.0)); + assert_eq!(model.output_tps, 20.0); + assert_eq!(model.queue_size, 2); + assert_eq!(model.queued_input_size, Some(17)); + assert_eq!(model.num_running_queries, 3); + assert_eq!(model.max_engine_concurrency, Some(8)); + assert_eq!(model.total_query_input_size, Some(31)); + assert_eq!(model.input_processing_queries, 1); + assert_eq!(model.output_generation_queries, 2); + assert_eq!(model.source_observed_at_unix_ms, 5); + assert!(model.complete); + } } #[test] - fn marks_inconsistent_totals_incomplete() { - let envelope = kv_snapshot_from_proto(snapshot(100, 40, 50)).unwrap(); - assert!(!envelope.models[0].complete); + fn unknown_exact_input_gauges_do_not_deactivate_the_model() { + let identity = pool(1); + let mut model = model_load("model-a", vec![identity.clone()]); + model.pending_first_output_input_tokens = None; + model.live_input_tokens = None; + let snapshot = proto::LoadSnapshot { + metadata: Some(metadata()), + window_ms: 1_000, + pools: vec![load_pool(identity)], + models: vec![model], + }; + + let translated = load_snapshot_from_proto(snapshot).unwrap(); + + assert_eq!(translated.relay_models.get("model-a"), Some(&true)); + for stats in translated.stats.models { + assert!(stats.complete); + assert_eq!(stats.queued_input_size, None); + assert_eq!(stats.total_query_input_size, None); + } } #[test] - fn rejects_duplicate_model_identity() { - let mut value = snapshot(100, 40, 60); - value.models.push(proto::ModelKvStats { - model: "alias-a".to_string(), - aliases: Vec::new(), - routing_cache: None, + fn model_without_a_frontend_or_serving_pool_remains_advertised_inactive() { + let identity = pool(1); + let mut relay_only = model_load("relay-only", vec![identity.clone()]); + relay_only.ready_frontends = None; + relay_only.pending_first_output_requests = None; + relay_only.pending_first_output_input_tokens = None; + relay_only.live_input_tokens = None; + relay_only.input_processing_requests = None; + relay_only.output_generation_requests = None; + relay_only.status = proto::DataStatus::Unavailable as i32; + relay_only.expected_frontends = 1; + relay_only.observed_frontends = 0; + + let snapshot = proto::LoadSnapshot { + metadata: Some(metadata()), + window_ms: 1_000, + pools: vec![load_pool(identity)], + models: vec![relay_only, model_load("frontend-only", Vec::new())], + }; + let translated = load_snapshot_from_proto(snapshot).unwrap(); + + assert_eq!(translated.relay_models.get("relay-only"), Some(&false)); + assert_eq!( + translated.relay_models.get("relay-only-alias"), + Some(&false) + ); + assert_eq!(translated.relay_models.get("frontend-only"), Some(&false)); + assert_eq!( + translated.relay_models.get("frontend-only-alias"), + Some(&false) + ); + } + + #[test] + fn malformed_or_unknown_pool_identity_rejects_the_snapshot() { + let mut malformed = pool(1); + malformed.cache_semantics_digest.pop(); + let malformed_snapshot = proto::KvUsageSnapshot { + metadata: Some(metadata()), + pools: vec![proto::PoolKvUsage { + pool: Some(malformed), + models: vec![registration("model-a")], + role: proto::WorkerRole::Aggregated as i32, + block_size_tokens: 1, + expected_ranks: 1, + observed_ranks: 1, + capacity_blocks: Some(1), + used_blocks: Some(0), + status: proto::DataStatus::Complete as i32, + source_observed_at_unix_ms: 1, + }], + }; + assert!(kv_snapshot_from_proto(malformed_snapshot).is_err()); + + let load_snapshot = proto::LoadSnapshot { + metadata: Some(metadata()), + window_ms: 1_000, pools: Vec::new(), - }); - assert!(kv_snapshot_from_proto(value).is_err()); + models: vec![model_load("model-a", vec![pool(9)])], + }; + assert!(load_snapshot_from_proto(load_snapshot).is_err()); } } diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs index a12bd81d5..f4bcb5f95 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/projection.rs @@ -101,7 +101,6 @@ impl StatsAggregator { let model_state = self.per_model.get_mut(&model_id)?; model_state.metrics.kv_cache = valid_kv_cache(&kv_cache).then_some(kv_cache); model_state.metrics.kv_cache_stats_observed = true; - model_state.metrics.kv_cache_received_at = Some(TokioInstant::now()); model_state.metrics.stats_observed_at_unix_ms = current_unix_millis(); let generation = model_state.generation.clone(); let stats = self.snapshot(&model_id); @@ -111,7 +110,6 @@ impl StatsAggregator { pub(super) fn apply_kv_cache_snapshot( &mut self, snapshot: KvCacheStatsEnvelope, - received_at: TokioInstant, ) -> Vec { let model_ids = self.per_model.keys().cloned().collect::>(); let mut changed = Vec::new(); @@ -125,12 +123,9 @@ impl StatsAggregator { .get_mut(&model_id) .expect("model id came from the current generation map"); if model_state.metrics.kv_cache == next { - model_state.metrics.kv_cache_received_at = next.as_ref().map(|_| received_at); continue; } model_state.metrics.kv_cache = next; - model_state.metrics.kv_cache_received_at = - model_state.metrics.kv_cache.as_ref().map(|_| received_at); model_state.metrics.kv_cache_stats_observed |= observed.is_some(); model_state.metrics.stats_observed_at_unix_ms = current_unix_millis(); changed.push(model_id); diff --git a/src/libraries/rust/stargate/crates/pylon/src/main.rs b/src/libraries/rust/stargate/crates/pylon/src/main.rs index eec83d7aa..7c104926e 100644 --- a/src/libraries/rust/stargate/crates/pylon/src/main.rs +++ b/src/libraries/rust/stargate/crates/pylon/src/main.rs @@ -23,6 +23,7 @@ use stargate_protocol::tunnel_contract::HEADER_STARGATE_UPSTREAM_RETRYABLE; const DEFAULT_PYLON_RETRYABLE_UPSTREAM_STATUS_CODES: &str = "429,503"; const DEFAULT_PYLON_UPSTREAM_RETRY_HEADER: &str = HEADER_STARGATE_UPSTREAM_RETRYABLE; const DEFAULT_OTEL_SERVICE_NAME: &str = "pylon"; +const DEFAULT_DYNAMO_RELAY_GRPC_URL: &str = "http://127.0.0.1:50051"; mod startup; @@ -32,6 +33,9 @@ struct Args { /// Base URL of the upstream HTTP inference server (for example http://127.0.0.1:8090) #[arg(long, value_name = "URL")] upstream_http_base_url: String, + /// KVDCRelay gRPC URL used for canonical load and KV-usage streams + #[arg(long, default_value = DEFAULT_DYNAMO_RELAY_GRPC_URL, value_name = "URL")] + dynamo_relay_grpc_url: String, /// QUIC tunnel listen address (advertised to stargate in forward mode) #[arg(long, default_value = "127.0.0.1:0", value_name = "ADDR")] quic_listen_addr: String, @@ -110,8 +114,8 @@ struct Args { /// Timeout for calibration requests in milliseconds #[arg(long, default_value_t = 30000, value_name = "MS")] bringup_calibration_timeout_ms: u64, - /// Engine stats stream source selection mode - #[arg(long, default_value_t = EngineStatsStreamMode::Auto, value_name = "MODE")] + /// Dynamo Relay aggregate stats stream mode + #[arg(long, default_value_t = EngineStatsStreamMode::Required, value_name = "MODE")] engine_stats_stream: EngineStatsStreamMode, /// Keep --initial-input-tps fixed for deterministic benchmark/test experiments #[arg(long, default_value_t = false, hide = true)] @@ -574,15 +578,12 @@ mod tests { } #[test] - fn engine_stats_stream_defaults_to_auto_mode() { + fn engine_stats_stream_is_required_by_default() { let args = parse_args(""); let metrics_config = stats_collector_config_from_args(&args); - assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Auto); - assert!( - !metrics_config.openai_fallback_stats_enabled, - "auto mode should wait for a permanent unsupported stream response before fallback stats" - ); + assert_eq!(args.engine_stats_stream, EngineStatsStreamMode::Required); + assert!(!metrics_config.openai_fallback_stats_enabled); } #[test] diff --git a/src/libraries/rust/stargate/crates/pylon/src/startup.rs b/src/libraries/rust/stargate/crates/pylon/src/startup.rs index c8a1cd3aa..941a4e83d 100644 --- a/src/libraries/rust/stargate/crates/pylon/src/startup.rs +++ b/src/libraries/rust/stargate/crates/pylon/src/startup.rs @@ -224,7 +224,7 @@ fn model_source_from_args(args: &Args) -> Result { struct RunningPylon { registration_client: InferenceServerRegistrationClient, - engine_stats_stream: Option, + engine_stats_stream: Option, stats_collector: StatsCollectorHandle, model_lifecycle: ModelLifecycleHandle, metrics_server: MetricsServerHandle, @@ -233,19 +233,13 @@ struct RunningPylon { initial_model_ids: Vec, } -struct RunningEngineStatsStream { - mode: EngineStatsStreamMode, - handle: EngineStatsStreamHandle, -} - impl RunningPylon { async fn run_until_shutdown(mut self, signal: S) -> Result<()> where S: Future>, { tokio::pin!(signal); - loop { - let error = tokio::select! { + let error = tokio::select! { result = signal.as_mut() => { let result = result.context("failed to receive pylon termination signal"); if let Ok(signal) = &result { @@ -257,20 +251,10 @@ impl RunningPylon { result = self.registration_client.wait_for_exit() => critical_task_exit_error("registration session", result), result = async { match self.engine_stats_stream.as_mut() { - Some(stream) => stream.handle.wait_for_exit().await, + Some(stream) => stream.wait_for_exit().await, None => std::future::pending().await, } - } => { - if engine_stats_exit_is_expected( - self.engine_stats_stream.as_ref().map(|stream| stream.mode), - &result, - ) { - info!("auto engine stats stream completed after enabling fallback"); - self.engine_stats_stream = None; - continue; - } - critical_task_exit_error("engine stats stream", result) - } + } => critical_task_exit_error("engine stats stream", result), result = self.stats_collector.wait_for_exit() => critical_task_exit_error("stats collector", result), result = async { self.model_lifecycle.wait_for_exit().await @@ -284,11 +268,10 @@ impl RunningPylon { } => { critical_task_exit_error("direct tunnel accept loop", result) } - }; - error!(error = %error, "critical pylon task exited"); - self.shutdown().await; - return Err(error); - } + }; + error!(error = %error, "critical pylon task exited"); + self.shutdown().await; + Err(error) } async fn shutdown(self) { @@ -305,7 +288,7 @@ impl RunningPylon { registration_client.shutdown(), async move { if let Some(stream) = engine_stats_stream { - stream.handle.shutdown().await; + stream.shutdown().await; } }, stats_collector.shutdown(), @@ -327,10 +310,6 @@ fn critical_task_exit_error(name: &'static str, result: TaskExit) -> anyhow::Err } } -fn engine_stats_exit_is_expected(mode: Option, result: &TaskExit) -> bool { - mode == Some(EngineStatsStreamMode::Auto) && result.is_ok() -} - async fn start_pylon_runtime(args: &Args, plan: &PylonStartupPlan) -> Result { let grpc_tls_ca_cert_pem = load_grpc_tls_ca_cert(args)?; let metrics = PylonMetrics::new()?; @@ -349,14 +328,9 @@ async fn start_pylon_runtime(args: &Args, plan: &PylonStartupPlan) -> Result Result, stats_config: &StatsCollectorConfig, runtime_state: PylonRuntimeState, ) -> Option<( - RunningEngineStatsStream, + EngineStatsStreamHandle, flume::Receiver, )> { let (stats_update_tx, stats_update_rx) = stats_aggregator_update_channel(stats_config); - let mut config = EngineStatsStreamConfig::new(&plan.upstream, args.engine_stats_stream); + let mut config = + EngineStatsStreamConfig::new(&args.dynamo_relay_grpc_url, args.engine_stats_stream); config.metrics = Some(metrics); config.runtime_state = Some(runtime_state); - let mode = config.mode; - start_engine_stats_stream(config, stats_update_tx) - .map(|handle| (RunningEngineStatsStream { mode, handle }, stats_update_rx)) + start_engine_stats_stream(config, stats_update_tx).map(|handle| (handle, stats_update_rx)) } async fn start_direct_tunnel_from_plan( @@ -724,8 +696,8 @@ mod tests { use axum::{Json, Router}; use clap::Parser; use pylon_lib::{ - EngineStatsStreamMode, PylonMetrics, RequestObservation, RequestObservationEndpoint, - RequestObservationState, TunnelTransportProtocol, + PylonMetrics, RequestObservation, RequestObservationEndpoint, RequestObservationState, + TunnelTransportProtocol, }; use stargate_proto::pb::stargate_control_plane_server::{ StargateControlPlane, StargateControlPlaneServer, @@ -1117,50 +1089,24 @@ mod tests { #[test] fn engine_stats_runtime_off_mode_leaves_stats_updates_unclaimed() { - let (args, plan) = startup(&["--engine-stats-stream", "off"]); + let (args, _) = startup(&["--engine-stats-stream", "off"]); let metrics = PylonMetrics::new().expect("metrics should initialize"); let config = StatsCollectorConfig::default(); assert!( - start_engine_stats_runtime( - &args, - &plan, - metrics, - &config, - PylonRuntimeState::default(), - ) - .is_none() + start_engine_stats_runtime(&args, metrics, &config, PylonRuntimeState::default(),) + .is_none() ); } #[tokio::test] async fn engine_stats_runtime_required_mode_claims_stats_updates() { - let (args, plan) = startup(&["--engine-stats-stream", "required"]); + let (args, _) = startup(&["--engine-stats-stream", "required"]); let metrics = PylonMetrics::new().expect("metrics should initialize"); let config = StatsCollectorConfig::default(); - let (engine_stats_stream, _stats_update_rx) = start_engine_stats_runtime( - &args, - &plan, - metrics, - &config, - PylonRuntimeState::default(), - ) - .expect("required engine stats should start a stream task"); - engine_stats_stream.handle.shutdown().await; - } - - #[test] - fn only_successful_auto_engine_stats_completion_is_nonfatal() { - let completed = Ok(()); - - assert!(engine_stats_exit_is_expected( - Some(EngineStatsStreamMode::Auto), - &completed, - )); - assert!(!engine_stats_exit_is_expected( - Some(EngineStatsStreamMode::Required), - &completed, - )); - assert!(!engine_stats_exit_is_expected(None, &completed)); + let (engine_stats_stream, _stats_update_rx) = + start_engine_stats_runtime(&args, metrics, &config, PylonRuntimeState::default()) + .expect("required engine stats should start a stream task"); + engine_stats_stream.shutdown().await; } #[tokio::test] diff --git a/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs b/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs index 2c3547bb0..28e0dd702 100644 --- a/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs +++ b/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs @@ -28,7 +28,7 @@ use crate::common::{ }; use axum::body::Body; use axum::extract::State; -use axum::http::{HeaderMap, Response}; +use axum::http::Response; use axum::routing::{get, post}; use axum::{Json, Router}; use futures::Stream; @@ -41,9 +41,9 @@ use pylon_lib::{ }; use stargate::routing::RoutingTargetKey; use stargate::test_support::StargateState; -use stargate_proto::{dynamo_frontend_stats as stats_proto, pb::InferenceServerStatus}; +use stargate_proto::{dynamo_kv_dc_relay as stats_proto, pb::InferenceServerStatus}; use tokio::net::TcpListener; -use tokio::sync::{broadcast, watch}; +use tokio::sync::watch; #[tokio::test] async fn end_to_end_registration_and_proxy() { @@ -496,7 +496,6 @@ async fn reverse_tunnel_handshake_rejects_non_reverse_instance_id() { #[derive(Clone)] struct EngineStatsState { model: String, - stats_tx: broadcast::Sender, connected_tx: watch::Sender, } @@ -511,18 +510,16 @@ async fn start_engine_stats_inst( ) { let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let addr = listener.local_addr().unwrap(); - let (stats_tx, _) = broadcast::channel(16); let (connected_tx, connected_rx) = watch::channel(false); let state = EngineStatsState { model: model.to_string(), - stats_tx, connected_tx, }; - let grpc = tonic::service::Routes::new( - stats_proto::frontend_stats_server::FrontendStatsServer::new(EngineStatsGrpc { + let grpc = tonic::service::Routes::new(stats_proto::kv_dc_relay_server::KvDcRelayServer::new( + EngineStatsGrpc { state: state.clone(), - }), - ) + }, + )) .into_axum_router(); let app = Router::new() .route("/v1/chat/completions", post(engine_stats_chat)) @@ -544,7 +541,6 @@ async fn start_engine_stats_inst( } async fn engine_stats_chat( - headers: HeaderMap, State(state): State, Json(req): Json, ) -> Response { @@ -555,13 +551,7 @@ async fn engine_stats_chat( .unwrap(); } - let request_id = headers - .get("x-request-id") - .and_then(|value| value.to_str().ok()) - .expect("test proxy should send x-request-id"); let model = state.model.clone(); - send_engine_stats_event(&state.stats_tx, request_id, &model, Some(1), Some(0), false); - send_engine_stats_event(&state.stats_tx, request_id, &model, Some(1), Some(2), true); let data_chunk = format!( r#"{{"object":"chat.completion.chunk","model":"{model}","choices":[{{"delta":{{"content":"Hello from engine stats"}}}}]}}"# @@ -578,71 +568,131 @@ data: [DONE]\n\n" .unwrap() } -fn send_engine_stats_event( - tx: &broadcast::Sender, - request_id: &str, - model: &str, - tokens_processed: Option, - tokens_generated: Option, - finished: bool, -) { - let _ = tx.send(stats_proto::StatsUpdate { - update: Some(stats_proto::stats_update::Update::RequestStats( - stats_proto::RequestStats { - request_id: request_id.to_string(), - model: model.to_string(), - tokens_processed, - tokens_generated, - finished, - }, - )), - }); -} - #[derive(Clone)] struct EngineStatsGrpc { state: EngineStatsState, } #[tonic::async_trait] -impl stats_proto::frontend_stats_server::FrontendStats for EngineStatsGrpc { - type WatchStatsStream = - Pin> + Send>>; - type WatchKvPlacementsStream = - Pin> + Send>>; +impl stats_proto::kv_dc_relay_server::KvDcRelay for EngineStatsGrpc { + type WatchKvCuckooFilterStream = Pin< + Box> + Send>, + >; + type WatchKvUsageStream = + Pin> + Send>>; + type WatchLoadStream = + Pin> + Send>>; + + async fn watch_kv_cuckoo_filter( + &self, + _request: tonic::Request<()>, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "the Pylon integration fixture does not model CKF data", + )) + } - async fn watch_stats( + async fn watch_kv_usage( &self, - _request: tonic::Request, - ) -> Result, tonic::Status> { - let _ = self.state.connected_tx.send(true); - let mut events = self.state.stats_tx.subscribe(); + _request: tonic::Request<()>, + ) -> Result, tonic::Status> { + let model = self.state.model.clone(); let stream = async_stream::stream! { - yield Ok(stats_proto::StatsUpdate { - update: Some(stats_proto::stats_update::Update::KvStats( - stats_proto::KvStatsSnapshot { - snapshot_id: 1, - observed_at_unix_ms: 1, - models: Vec::new(), - }, - )), - }); loop { - match events.recv().await { - Ok(event) => yield Ok(event), - Err(broadcast::error::RecvError::Lagged(_)) => break, - Err(broadcast::error::RecvError::Closed) => break, - } + yield Ok(stats_proto::KvUsageSnapshot { + metadata: Some(relay_metadata()), + pools: vec![stats_proto::PoolKvUsage { + pool: Some(stats_pool_identity()), + models: vec![stats_model_registration(&model)], + role: stats_proto::WorkerRole::Aggregated as i32, + block_size_tokens: 1, + expected_ranks: 1, + observed_ranks: 1, + capacity_blocks: Some(1_000), + used_blocks: Some(400), + status: stats_proto::DataStatus::Complete as i32, + source_observed_at_unix_ms: 1, + }], + }); + tokio::time::sleep(Duration::from_millis(100)).await; } }; Ok(tonic::Response::new(Box::pin(stream))) } - async fn watch_kv_placements( + async fn watch_load( &self, - _request: tonic::Request, - ) -> Result, tonic::Status> { - Err(tonic::Status::unimplemented("placements are not used here")) + _request: tonic::Request<()>, + ) -> Result, tonic::Status> { + let _ = self.state.connected_tx.send(true); + let model = self.state.model.clone(); + let stream = async_stream::stream! { + loop { + yield Ok(stats_proto::LoadSnapshot { + metadata: Some(relay_metadata()), + window_ms: 1_000, + pools: vec![stats_proto::PoolLoad { + pool: Some(stats_pool_identity()), + role: stats_proto::WorkerRole::Aggregated as i32, + live_workers: Some(1), + active_prefill_tokens: Some(17), + active_decode_blocks: Some(3), + max_concurrency: Some(8), + scheduler_status: stats_proto::DataStatus::Complete as i32, + scheduler_observed_at_unix_ms: 1, + }], + models: vec![stats_proto::ModelLoad { + model: Some(stats_model_registration(&model)), + ready_frontends: Some(1), + pending_first_output_requests: Some(2), + pending_first_output_input_tokens: Some(17), + live_input_tokens: Some(31), + input_processing_requests: Some(1), + output_generation_requests: Some(2), + serving_pools: vec![stats_pool_identity()], + requests_started: 4, + requests_completed: 1, + requests_failed: 0, + requests_cancelled: 0, + input_tokens: Some(31), + output_tokens: 20, + status: stats_proto::DataStatus::Complete as i32, + expected_frontends: 1, + observed_frontends: 1, + source_observed_at_unix_ms: 1, + }], + }); + tokio::time::sleep(Duration::from_millis(100)).await; + } + }; + Ok(tonic::Response::new(Box::pin(stream))) + } +} + +fn relay_metadata() -> stats_proto::RelayMessageMetadata { + stats_proto::RelayMessageMetadata { + drt_instance_id: 1, + relay_incarnation: 1, + observed_at_unix_ms: 1, + } +} + +fn stats_pool_identity() -> stats_proto::PoolIdentity { + stats_proto::PoolIdentity { + cache_semantics_digest: vec![1; 16], + cache_semantics_source: stats_proto::IdentitySource::DefaultDerived as i32, + routing_scope_digest: vec![2; 16], + routing_scope_source: stats_proto::IdentitySource::DefaultDerived as i32, + dc_id: 1, + } +} + +fn stats_model_registration(model: &str) -> stats_proto::ModelRegistration { + stats_proto::ModelRegistration { + model: model.to_string(), + base_model: model.to_string(), + adapter: None, + aliases: Vec::new(), } } @@ -685,16 +735,26 @@ async fn wait_for_engine_stats_stream_stats( loop { let candidates = state.candidates_for_target(&target).await; if candidates.iter().any(|candidate| { - candidate - .stats - .stats_capabilities - .iter() - .any(|capability| capability == "model.throughput.engine_stream") - && candidate - .stats + let stats = &candidate.stats; + stats.output_tps == 20.0 + && stats.queue_size == 2 + && stats.queued_input_size == 17 + && stats.num_running_queries == 3 + && stats.max_engine_concurrency == 8 + && stats.total_query_input_size == 31 + && stats.input_processing_queries == 1 + && stats.output_generation_queries == 2 + && stats.kv_cache_capacity_tokens == 1_000 + && stats.kv_cache_used_tokens == 400 + && stats.kv_cache_free_tokens == 600 + && stats + .stats_sources + .iter() + .any(|source| source == "dynamo_relay_load") + && stats .stats_sources .iter() - .any(|source| source == "engine_stats_stream") + .any(|source| source == "dynamo_relay_kv_usage") }) { return; } @@ -712,7 +772,7 @@ async fn wait_for_engine_stats_stream_stats( }) .collect::>(); panic!( - "model '{model_id}' did not report engine stats stream stats within {}s; last_seen={last_seen:?}", + "model '{model_id}' did not report canonical Relay stats within {}s; last_seen={last_seen:?}", timeout.as_secs() ); } From 5617d4972d6eaef7a798451da716faf17c231696 Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Thu, 27 Aug 2026 20:49:32 +0000 Subject: [PATCH 8/9] fix(pylon): preserve passthrough tunnel semantics Signed-off-by: Barry Greengus --- .../pylon-lib/src/quic_http_tunnel/backend.rs | 16 ++++++--- .../pylon-lib/src/quic_http_tunnel/core.rs | 10 ++++-- .../pylon-lib/src/quic_http_tunnel/tests.rs | 34 ++++++++++++++----- .../crates/pylon-lib/src/runtime_state.rs | 8 +---- 4 files changed, 45 insertions(+), 23 deletions(-) diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs index 070f3faa0..4ff915567 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/backend.rs @@ -69,20 +69,26 @@ pub(crate) mod dynamo { pub(crate) const HEADER_REQUEST_PRIORITY: &str = "x-dynamo-request-priority"; pub(crate) const HEADER_REQUEST_STRICT_PRIORITY: &str = "x-dynamo-request-strict-priority"; - /// Platform metadata is consumed by pylon and never becomes part of the + /// Engine headers owned by pylon are never accepted from callers. + const STRIPPED_ENGINE_HEADERS: [&str; 2] = + [HEADER_REQUEST_PRIORITY, HEADER_REQUEST_STRICT_PRIORITY]; + + /// Platform metadata is consumed by pylon and never becomes part of a /// Dynamo request. Priority is translated below; request state is local. - const STRIPPED_REQUEST_HEADERS: [&str; 7] = [ + const PLATFORM_METADATA_HEADERS: [&str; 5] = [ "request-id", "x-dynamo-request-id", HEADER_MODEL, HEADER_ROUTING_KEY, HEADER_PRIORITY, - HEADER_REQUEST_PRIORITY, - HEADER_REQUEST_STRICT_PRIORITY, ]; pub(crate) fn is_stripped_engine_header(name: &HeaderName) -> bool { - STRIPPED_REQUEST_HEADERS.contains(&name.as_str()) + STRIPPED_ENGINE_HEADERS.contains(&name.as_str()) + } + + pub(crate) fn is_platform_metadata_header(name: &HeaderName) -> bool { + PLATFORM_METADATA_HEADERS.contains(&name.as_str()) } /// Map the platform rank (lower wins, absent = unconfigured) to the diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs index 4403b189f..7ea48d8df 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/core.rs @@ -757,7 +757,7 @@ async fn send_upstream_request( }; let mut upstream_headers = HeaderMap::with_capacity(request_headers.len()); for (name, value) in request_headers { - if should_forward_header(name, &app.retry) { + if should_forward_header(name, &app.retry, app.upstream_backend) { upstream_headers.append(name, value.clone()); } } @@ -1117,9 +1117,15 @@ pub(super) fn join_base_path(base: &str, path_and_query: &str) -> Result bool { +pub(super) fn should_forward_header( + name: &HeaderName, + retry: &PylonRetryConfig, + upstream_backend: UpstreamBackend, +) -> bool { !is_tunnel_control_header(name, retry) && !backend::dynamo::is_stripped_engine_header(name) + && (upstream_backend != UpstreamBackend::Dynamo + || !backend::dynamo::is_platform_metadata_header(name)) && !matches!( name.as_str(), "host" | "x-method" | "x-path" | HEADER_STARGATE_EXPECTED_QUEUE_MS diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs index 3aec2e3bc..f627f046e 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/quic_http_tunnel/tests.rs @@ -480,8 +480,6 @@ fn pylon_request_header_filter_strips_tunnel_headers_case_insensitively() "X-Method", "X-Path", "X-Stargate-Expected-Queue-Ms", - "Request-Id", - "X-Dynamo-Request-Id", "X-Dynamo-Request-Priority", "X-Dynamo-Request-Strict-Priority", ] @@ -490,12 +488,33 @@ fn pylon_request_header_filter_strips_tunnel_headers_case_insensitively() { assert!(!should_forward_header( &HeaderName::from_bytes(name.as_bytes())?, - &retry + &retry, + UpstreamBackend::Passthrough, + )); + } + for name in [ + "Request-Id", + "X-Dynamo-Request-Id", + "X-Model", + "X-Routing-Key", + "X-Priority", + ] { + let name = HeaderName::from_bytes(name.as_bytes())?; + assert!(should_forward_header( + &name, + &retry, + UpstreamBackend::Passthrough, + )); + assert!(!should_forward_header( + &name, + &retry, + UpstreamBackend::Dynamo, )); } assert!(should_forward_header( &HeaderName::from_bytes(b"X-Request-Id")?, - &retry + &retry, + UpstreamBackend::Dynamo, )); Ok(()) } @@ -549,12 +568,12 @@ fn pylon_consumes_platform_metadata_instead_of_forwarding_it_to_dynamo() { "x-routing-key", "x-priority", ] { - assert!(dynamo::is_stripped_engine_header( + assert!(dynamo::is_platform_metadata_header( &HeaderName::from_bytes(name.as_bytes()).unwrap() )); } for name in ["x-request-id", "x-input-tokens"] { - assert!(!dynamo::is_stripped_engine_header( + assert!(!dynamo::is_platform_metadata_header( &HeaderName::from_bytes(name.as_bytes()).unwrap() )); } @@ -1070,7 +1089,6 @@ async fn http3_direct_tunnel_accepts_responses_request_to_upstream() { ); let mut config = test_tunnel_config_for(app).await; config.tunnel_protocol = TunnelTransportProtocol::Http3; - config.forwarding.upstream_backend = UpstreamBackend::Passthrough; let tunnel = start_quic_http_tunnel(config).await.unwrap(); let mut headers = HeaderMap::new(); headers.insert("x-request-id", "req-h3-direct".parse().unwrap()); @@ -1676,7 +1694,6 @@ async fn quic_tunnel_forwards_to_http_backend() { }), ); let (mut config, _metrics) = metered_test_tunnel_config_for(app).await; - config.forwarding.upstream_backend = UpstreamBackend::Passthrough; config.forwarding.retry.upstream_retry_header = HeaderName::from_static("x-vendor-retryable"); let mut tunnel = RawTunnelTest::start(config).await; @@ -2768,7 +2785,6 @@ async fn assert_direct_embeddings_case(case: DirectEmbeddingsCase) { let (runtime_state, rx) = observed_runtime(16); let mut config = test_tunnel_config_for(app).await; config.tunnel_protocol = case.protocol; - config.forwarding.upstream_backend = UpstreamBackend::Passthrough; config.forwarding.runtime_state = runtime_state; let tunnel = start_quic_http_tunnel(config).await.unwrap(); let client = DirectTunnelClient::connect(case.protocol, tunnel.listen_addr()).await; diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs index 63b45b1e8..4177ea97f 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/runtime_state.rs @@ -90,7 +90,7 @@ impl ModelGeneration { } } -#[derive(Clone, Debug)] +#[derive(Clone, Debug, Default)] pub struct PylonRuntimeState { advertised: Arc>, live_requests: LiveRequestState, @@ -623,12 +623,6 @@ impl PylonRuntimeState { } } -impl Default for PylonRuntimeState { - fn default() -> Self { - Self::new(InferenceServerStatus::Unknown, &[]) - } -} - impl RequestObservationEvent { pub fn observation(&self) -> &RequestObservation { &self.observation From b08b23ed014083c59f60846bf87188ddc0f3f132 Mon Sep 17 00:00:00 2001 From: Barry Greengus Date: Fri, 28 Aug 2026 18:29:05 +0000 Subject: [PATCH 9/9] fix(pylon): align Dynamo relay stats contract Signed-off-by: Barry Greengus --- dependencies.md | 1 - .../crates/mock-dynamo/src/stats_stream.rs | 56 +-- .../rust/stargate/crates/proto/build.rs | 15 +- .../proto/proto/dynamo_kv_dc_relay.proto | 17 +- .../pylon-lib/src/registration/tests.rs | 3 +- .../src/stats/engine_stats_stream.rs | 10 +- .../crates/pylon-lib/src/stats/kv_stats.rs | 426 ++++++++++++------ .../stargate/tests/suite/integration.rs | 17 +- 8 files changed, 354 insertions(+), 191 deletions(-) diff --git a/dependencies.md b/dependencies.md index 24b448570..045e5b337 100644 --- a/dependencies.md +++ b/dependencies.md @@ -912,7 +912,6 @@ Generated by `go run -C ./tools/collect-dependencies .`. Refresh: `go run -C ./t ## Apache-2.0 OR MIT -- `Rust`: `criterion 0.5` - `Rust`: `indexmap =2.14.0` - `Rust`: `pin-project-lite 0.2` - `Rust`: `uom 0.36` diff --git a/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs b/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs index 9f77ede93..ca3ad2e39 100644 --- a/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs +++ b/src/libraries/rust/stargate/crates/mock-dynamo/src/stats_stream.rs @@ -120,7 +120,7 @@ fn snapshot_interval() -> tokio::time::Interval { #[derive(Default)] struct LoadAccumulator { live: HashMap, - windows: BTreeMap, + totals: BTreeMap, } struct LiveRequest { @@ -130,7 +130,7 @@ struct LiveRequest { } #[derive(Default)] -struct WindowCounters { +struct TrafficCounters { requests_started: u64, requests_completed: u64, input_tokens: u64, @@ -139,20 +139,20 @@ struct WindowCounters { impl LoadAccumulator { fn observe(&mut self, event: StatsStreamEvent) { - let window = self.windows.entry(event.model.clone()).or_default(); + let totals = self.totals.entry(event.model.clone()).or_default(); if let Some(request) = self.live.get_mut(&event.request_id) { - window.input_tokens = window + totals.input_tokens = totals .input_tokens .saturating_add(event.input_tokens.saturating_sub(request.input_tokens)); - window.output_tokens = window + totals.output_tokens = totals .output_tokens .saturating_add(event.output_tokens.saturating_sub(request.output_tokens)); request.input_tokens = request.input_tokens.max(event.input_tokens); request.output_tokens = request.output_tokens.max(event.output_tokens); } else { - window.requests_started = window.requests_started.saturating_add(1); - window.input_tokens = window.input_tokens.saturating_add(event.input_tokens); - window.output_tokens = window.output_tokens.saturating_add(event.output_tokens); + totals.requests_started = totals.requests_started.saturating_add(1); + totals.input_tokens = totals.input_tokens.saturating_add(event.input_tokens); + totals.output_tokens = totals.output_tokens.saturating_add(event.output_tokens); self.live.insert( event.request_id.clone(), LiveRequest { @@ -164,19 +164,18 @@ impl LoadAccumulator { } if event.finished { self.live.remove(&event.request_id); - window.requests_completed = window.requests_completed.saturating_add(1); + totals.requests_completed = totals.requests_completed.saturating_add(1); } } - fn snapshot(&mut self, configured_model: &str) -> proto::LoadSnapshot { + fn snapshot(&self, configured_model: &str) -> proto::LoadSnapshot { let mut model_ids = BTreeSet::from([configured_model.to_string()]); model_ids.extend(self.live.values().map(|request| request.model.clone())); - model_ids.extend(self.windows.keys().cloned()); - let windows = std::mem::take(&mut self.windows); + model_ids.extend(self.totals.keys().cloned()); let models = model_ids .into_iter() .map(|model| { - let window = windows.get(&model); + let totals = self.totals.get(&model); let live = self .live .values() @@ -202,12 +201,16 @@ impl LoadAccumulator { input_processing_requests: Some(input_processing), output_generation_requests: Some(output_generation), serving_pools: vec![pool_identity()], - requests_started: window.map_or(0, |window| window.requests_started), - requests_completed: window.map_or(0, |window| window.requests_completed), - requests_failed: 0, - requests_cancelled: 0, - input_tokens: Some(window.map_or(0, |window| window.input_tokens)), - output_tokens: window.map_or(0, |window| window.output_tokens), + requests_started_total: Some( + totals.map_or(0, |totals| totals.requests_started), + ), + requests_completed_total: Some( + totals.map_or(0, |totals| totals.requests_completed), + ), + requests_failed_total: Some(0), + requests_cancelled_total: Some(0), + input_tokens_total: Some(totals.map_or(0, |totals| totals.input_tokens)), + output_tokens_total: Some(totals.map_or(0, |totals| totals.output_tokens)), status: proto::DataStatus::Complete as i32, expected_frontends: 1, observed_frontends: 1, @@ -217,7 +220,6 @@ impl LoadAccumulator { .collect(); proto::LoadSnapshot { metadata: Some(metadata()), - window_ms: 1_000, pools: vec![proto::PoolLoad { pool: Some(pool_identity()), role: proto::WorkerRole::Aggregated as i32, @@ -284,7 +286,7 @@ mod tests { use super::*; #[test] - fn load_snapshots_replace_window_counters_and_keep_live_gauges() { + fn load_snapshots_keep_cumulative_counters_and_live_gauges() { let mut accumulator = LoadAccumulator::default(); accumulator.observe(StatsStreamEvent { request_id: "req-1".to_string(), @@ -295,15 +297,15 @@ mod tests { }); let first = accumulator.snapshot("model-a"); - assert_eq!(first.models[0].requests_started, 1); - assert_eq!(first.models[0].input_tokens, Some(10)); - assert_eq!(first.models[0].output_tokens, 2); + assert_eq!(first.models[0].requests_started_total, Some(1)); + assert_eq!(first.models[0].input_tokens_total, Some(10)); + assert_eq!(first.models[0].output_tokens_total, Some(2)); assert_eq!(first.models[0].output_generation_requests, Some(1)); let second = accumulator.snapshot("model-a"); - assert_eq!(second.models[0].requests_started, 0); - assert_eq!(second.models[0].input_tokens, Some(0)); - assert_eq!(second.models[0].output_tokens, 0); + assert_eq!(second.models[0].requests_started_total, Some(1)); + assert_eq!(second.models[0].input_tokens_total, Some(10)); + assert_eq!(second.models[0].output_tokens_total, Some(2)); assert_eq!(second.models[0].output_generation_requests, Some(1)); } } diff --git a/src/libraries/rust/stargate/crates/proto/build.rs b/src/libraries/rust/stargate/crates/proto/build.rs index c22e9a5f5..3fa37b35c 100644 --- a/src/libraries/rust/stargate/crates/proto/build.rs +++ b/src/libraries/rust/stargate/crates/proto/build.rs @@ -31,9 +31,22 @@ fn compile_proto_plan(plan: ProtoCompilePlan) -> Result<(), Box>(); + let mut includes = plan + .includes + .iter() + .map(std::path::PathBuf::from) + .collect::>(); + if let Some(well_known_types) = std::env::var_os("PROTOC_WKT_DIR") { + includes.push(well_known_types.into()); + } builder .build_server(plan.build_server) - .compile_protos(plan.protos, plan.includes)?; + .compile_protos(&protos, &includes)?; Ok(()) } diff --git a/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto b/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto index 8f66678bb..031c99c24 100644 --- a/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto +++ b/src/libraries/rust/stargate/crates/proto/proto/dynamo_kv_dc_relay.proto @@ -159,9 +159,8 @@ message PoolKvUsage message LoadSnapshot { RelayMessageMetadata metadata = 1; - uint32 window_ms = 2; - repeated PoolLoad pools = 3; - repeated ModelLoad models = 4; + repeated PoolLoad pools = 2; + repeated ModelLoad models = 3; } message PoolLoad @@ -187,12 +186,12 @@ message ModelLoad optional uint64 output_generation_requests = 7; repeated PoolIdentity serving_pools = 8; - uint64 requests_started = 10; - uint64 requests_completed = 11; - uint64 requests_failed = 12; - uint64 requests_cancelled = 13; - optional uint64 input_tokens = 14; - uint64 output_tokens = 15; + optional uint64 requests_started_total = 10; + optional uint64 requests_completed_total = 11; + optional uint64 requests_failed_total = 12; + optional uint64 requests_cancelled_total = 13; + optional uint64 input_tokens_total = 14; + optional uint64 output_tokens_total = 15; DataStatus status = 20; uint32 expected_frontends = 21; diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs index 9f0df65e5..9486c96c7 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/registration/tests.rs @@ -538,8 +538,7 @@ fn stargate_grpc_endpoint_rejects_custom_ca_for_plaintext_http() { let error = endpoint .channel_endpoint(Some(b"private CA contents must not be logged")) - .err() - .expect("custom CA with plaintext HTTP should be rejected"); + .expect_err("custom CA with plaintext HTTP should be rejected"); assert!( error diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs index d44e57ffa..62bc7327f 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/engine_stats_stream.rs @@ -11,7 +11,7 @@ use stargate_runtime::OwnedTask; use tokio_util::sync::CancellationToken; use super::collector::StatsAggregatorUpdate; -use super::kv_stats::{kv_snapshot_from_proto, load_snapshot_from_proto}; +use super::kv_stats::{RelayLoadTranslator, kv_snapshot_from_proto}; use super::metrics::PylonMetrics; use crate::PylonRuntimeState; @@ -139,6 +139,8 @@ async fn run_load_stream( stop: CancellationToken, ) { let mut backoff = config.initial_reconnect_backoff; + let mut last_identity = None; + let mut translator = RelayLoadTranslator::default(); loop { if stop.is_cancelled() { return; @@ -201,6 +203,10 @@ async fn run_load_stream( epoch.load = Some(identity); changed }; + if last_identity.is_some_and(|current| current != identity) { + translator = RelayLoadTranslator::default(); + } + last_identity = Some(identity); if identity_changed { send_update( &updates, @@ -209,7 +215,7 @@ async fn run_load_stream( ) .await; } - match load_snapshot_from_proto(snapshot) { + match translator.translate(snapshot) { Ok(translation) => { if let Some(runtime_state) = &config.runtime_state { runtime_state.replace_relay_models(translation.relay_models); diff --git a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs index 28eb41fcc..c0b6680c1 100644 --- a/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs +++ b/src/libraries/rust/stargate/crates/pylon-lib/src/stats/kv_stats.rs @@ -43,6 +43,18 @@ pub(super) struct RelayLoadTranslation { pub(super) relay_models: BTreeMap, } +#[derive(Clone, Copy)] +struct CounterBaseline { + observed_at_unix_ms: u64, + input_tokens_total: Option, + output_tokens_total: u64, +} + +#[derive(Default)] +pub(super) struct RelayLoadTranslator { + counters: HashMap, +} + pub(super) fn kv_snapshot_from_proto( snapshot: proto::KvUsageSnapshot, ) -> anyhow::Result { @@ -143,136 +155,180 @@ pub(super) fn kv_snapshot_from_proto( Ok(KvCacheStatsEnvelope { models }) } -pub(super) fn load_snapshot_from_proto( - snapshot: proto::LoadSnapshot, -) -> anyhow::Result { - anyhow::ensure!(snapshot.window_ms > 0, "load snapshot window is zero"); - snapshot - .metadata - .as_ref() - .ok_or_else(|| anyhow::anyhow!("load snapshot metadata is missing"))?; - let window_seconds = f64::from(snapshot.window_ms) / 1_000.0; - let mut pools = HashMap::::new(); - for pool in snapshot.pools { - let key = pool_key(pool.pool)?; - let role = worker_role(pool.role)?; - anyhow::ensure!( - pools - .insert( - key, - LoadPool { - role, - live_workers: pool.live_workers, - max_concurrency: pool.max_concurrency, - complete: data_complete(pool.scheduler_status), - }, - ) - .is_none(), - "duplicate load pool" - ); - } - - let mut identity_owner = HashMap::::new(); - let mut relay_models = BTreeMap::new(); - let mut models = Vec::new(); - for model in snapshot.models { - let registration = model - .model +impl RelayLoadTranslator { + pub(super) fn translate( + &mut self, + snapshot: proto::LoadSnapshot, + ) -> anyhow::Result { + snapshot + .metadata .as_ref() - .ok_or_else(|| anyhow::anyhow!("load model registration is missing"))?; - let identities = - registration_identities(std::slice::from_ref(registration), &mut identity_owner)?; - let serving_pools = model - .serving_pools - .into_iter() - .map(|pool| pool_key(Some(pool))) - .collect::>>()?; - anyhow::ensure!( - serving_pools.iter().all(|pool| pools.contains_key(pool)), - "load model references an unknown serving pool" - ); - let selected_role = if serving_pools.iter().any(|pool| { - pools - .get(pool) - .is_some_and(|pool| pool.role == proto::WorkerRole::Aggregated) - }) { - proto::WorkerRole::Aggregated - } else { - proto::WorkerRole::Decode - }; - let selected = serving_pools - .iter() - .filter_map(|key| pools.get(key)) - .filter(|pool| pool.role == selected_role) - .collect::>(); - let scheduler_live = selected - .iter() - .any(|pool| pool.complete && pool.live_workers.is_some_and(|workers| workers > 0)); - let max_engine_concurrency = (!selected.is_empty() - && selected + .ok_or_else(|| anyhow::anyhow!("load snapshot metadata is missing"))?; + let mut pools = HashMap::::new(); + for pool in snapshot.pools { + let key = pool_key(pool.pool)?; + let role = worker_role(pool.role)?; + anyhow::ensure!( + pools + .insert( + key, + LoadPool { + role, + live_workers: pool.live_workers, + max_concurrency: pool.max_concurrency, + complete: data_complete(pool.scheduler_status), + }, + ) + .is_none(), + "duplicate load pool" + ); + } + + let mut identity_owner = HashMap::::new(); + let mut seen_models = HashSet::new(); + let mut next_counters = HashMap::new(); + let mut relay_models = BTreeMap::new(); + let mut models = Vec::new(); + for model in snapshot.models { + let registration = model + .model + .as_ref() + .ok_or_else(|| anyhow::anyhow!("load model registration is missing"))?; + let identities = + registration_identities(std::slice::from_ref(registration), &mut identity_owner)?; + let canonical_model = registration.model.trim(); + anyhow::ensure!( + seen_models.insert(canonical_model.to_string()), + "duplicate load model" + ); + let serving_pools = model + .serving_pools + .into_iter() + .map(|pool| pool_key(Some(pool))) + .collect::>>()?; + anyhow::ensure!( + serving_pools.iter().all(|pool| pools.contains_key(pool)), + "load model references an unknown serving pool" + ); + let selected_role = if serving_pools.iter().any(|pool| { + pools + .get(pool) + .is_some_and(|pool| pool.role == proto::WorkerRole::Aggregated) + }) { + proto::WorkerRole::Aggregated + } else { + proto::WorkerRole::Decode + }; + let selected = serving_pools + .iter() + .filter_map(|key| pools.get(key)) + .filter(|pool| pool.role == selected_role) + .collect::>(); + let scheduler_live = selected .iter() - .all(|pool| pool.complete && pool.max_concurrency.is_some())) - .then(|| { - selected.iter().try_fold(0_u64, |total, pool| { - total.checked_add(pool.max_concurrency?) + .any(|pool| pool.complete && pool.live_workers.is_some_and(|workers| workers > 0)); + let max_engine_concurrency = (!selected.is_empty() + && selected + .iter() + .all(|pool| pool.complete && pool.max_concurrency.is_some())) + .then(|| { + selected.iter().try_fold(0_u64, |total, pool| { + total.checked_add(pool.max_concurrency?) + }) }) - }) - .flatten(); - let required = ( - model.ready_frontends, - model.pending_first_output_requests, - model.input_processing_requests, - model.output_generation_requests, - ); - let complete = data_complete(model.status) - && model.expected_frontends > 0 - && model.observed_frontends == model.expected_frontends - && model.source_observed_at_unix_ms > 0 - && matches!(required, (Some(_), Some(_), Some(_), Some(_))); - let (ready_frontends, queue_size, input_processing_queries, output_generation_queries) = - required; - let num_running_queries = input_processing_queries - .zip(output_generation_queries) - .and_then(|(input, output)| input.checked_add(output)); - let complete = complete && num_running_queries.is_some(); - let active = complete - && !serving_pools.is_empty() - && ready_frontends.is_some_and(|ready| ready > 0) - && scheduler_live; - for identity in &identities { - relay_models.insert(identity.clone(), active); - models.push(RelayLoadStatsSnapshot { - model: identity.clone(), - input_tps: if complete { - model - .input_tokens - .map(|tokens| tokens as f64 / window_seconds) - } else { - None - }, - output_tps: if complete { - model.output_tokens as f64 / window_seconds - } else { - 0.0 - }, - queue_size: queue_size.unwrap_or_default(), - queued_input_size: complete - .then_some(model.pending_first_output_input_tokens) - .flatten(), - num_running_queries: num_running_queries.unwrap_or_default(), - max_engine_concurrency, - total_query_input_size: complete.then_some(model.live_input_tokens).flatten(), - input_processing_queries: input_processing_queries.unwrap_or_default(), - output_generation_queries: output_generation_queries.unwrap_or_default(), - source_observed_at_unix_ms: model.source_observed_at_unix_ms, - complete, - }); + .flatten(); + let required = ( + model.ready_frontends, + model.pending_first_output_requests, + model.input_processing_requests, + model.output_generation_requests, + ); + let complete = data_complete(model.status) + && model.expected_frontends > 0 + && model.observed_frontends == model.expected_frontends + && model.source_observed_at_unix_ms > 0 + && matches!(required, (Some(_), Some(_), Some(_), Some(_))); + let (ready_frontends, queue_size, input_processing_queries, output_generation_queries) = + required; + let num_running_queries = input_processing_queries + .zip(output_generation_queries) + .and_then(|(input, output)| input.checked_add(output)); + let load_complete = complete && num_running_queries.is_some(); + let rates = model + .output_tokens_total + .filter(|_| load_complete) + .map(|output_tokens_total| CounterBaseline { + observed_at_unix_ms: model.source_observed_at_unix_ms, + input_tokens_total: model.input_tokens_total, + output_tokens_total, + }) + .map(|current| { + let rates = self + .counters + .get(canonical_model) + .and_then(|previous| counter_rates(*previous, current)); + (current, rates) + }); + if let Some((current, _)) = rates { + next_counters.insert(canonical_model.to_string(), current); + } + let (input_tps, output_tps) = rates + .and_then(|(_, rates)| rates) + .map_or((None, None), |rates| rates); + let stats_complete = load_complete && output_tps.is_some(); + let active = load_complete + && !serving_pools.is_empty() + && ready_frontends.is_some_and(|ready| ready > 0) + && scheduler_live; + for identity in &identities { + relay_models.insert(identity.clone(), active); + models.push(RelayLoadStatsSnapshot { + model: identity.clone(), + input_tps: stats_complete.then_some(input_tps).flatten(), + output_tps: output_tps.unwrap_or_default(), + queue_size: queue_size.unwrap_or_default(), + queued_input_size: load_complete + .then_some(model.pending_first_output_input_tokens) + .flatten(), + num_running_queries: num_running_queries.unwrap_or_default(), + max_engine_concurrency, + total_query_input_size: load_complete + .then_some(model.live_input_tokens) + .flatten(), + input_processing_queries: input_processing_queries.unwrap_or_default(), + output_generation_queries: output_generation_queries.unwrap_or_default(), + source_observed_at_unix_ms: model.source_observed_at_unix_ms, + complete: stats_complete, + }); + } } + self.counters = next_counters; + Ok(RelayLoadTranslation { + stats: RelayLoadStatsEnvelope { models }, + relay_models, + }) } - Ok(RelayLoadTranslation { - stats: RelayLoadStatsEnvelope { models }, - relay_models, - }) +} + +fn counter_rates( + previous: CounterBaseline, + current: CounterBaseline, +) -> Option<(Option, Option)> { + let elapsed_ms = current + .observed_at_unix_ms + .checked_sub(previous.observed_at_unix_ms) + .filter(|elapsed| *elapsed > 0)?; + let elapsed_seconds = elapsed_ms as f64 / 1_000.0; + let input_tps = previous + .input_tokens_total + .zip(current.input_tokens_total) + .and_then(|(previous, current)| current.checked_sub(previous)) + .map(|tokens| tokens as f64 / elapsed_seconds); + let output_tps = current + .output_tokens_total + .checked_sub(previous.output_tokens_total) + .map(|tokens| tokens as f64 / elapsed_seconds); + Some((input_tps, output_tps)) } fn pool_key(pool: Option) -> anyhow::Result { @@ -414,16 +470,16 @@ mod tests { input_processing_requests: Some(1), output_generation_requests: Some(2), serving_pools, - requests_started: 4, - requests_completed: 3, - requests_failed: 0, - requests_cancelled: 0, - input_tokens: Some(40), - output_tokens: 20, + requests_started_total: Some(4), + requests_completed_total: Some(3), + requests_failed_total: Some(0), + requests_cancelled_total: Some(0), + input_tokens_total: Some(40), + output_tokens_total: Some(20), status: proto::DataStatus::Complete as i32, expected_frontends: 1, observed_frontends: 1, - source_observed_at_unix_ms: 5, + source_observed_at_unix_ms: 5_000, } } @@ -476,14 +532,27 @@ mod tests { #[test] fn complete_load_activates_model_and_alias() { let identity = pool(1); + let mut translator = RelayLoadTranslator::default(); + let mut baseline = model_load("model-a", vec![identity.clone()]); + baseline.input_tokens_total = Some(0); + baseline.output_tokens_total = Some(0); + baseline.source_observed_at_unix_ms = 4_000; + let first = translator + .translate(proto::LoadSnapshot { + metadata: Some(metadata()), + pools: vec![load_pool(identity.clone())], + models: vec![baseline], + }) + .unwrap(); + assert_eq!(first.relay_models.get("model-a"), Some(&true)); + assert!(first.stats.models.iter().all(|model| !model.complete)); let snapshot = proto::LoadSnapshot { metadata: Some(metadata()), - window_ms: 1_000, pools: vec![load_pool(identity.clone())], models: vec![model_load("model-a", vec![identity])], }; - let translated = load_snapshot_from_proto(snapshot).unwrap(); + let translated = translator.translate(snapshot).unwrap(); assert_eq!(translated.relay_models.get("model-a"), Some(&true)); assert_eq!(translated.relay_models.get("model-a-alias"), Some(&true)); assert_eq!(translated.stats.models.len(), 2); @@ -497,25 +566,98 @@ mod tests { assert_eq!(model.total_query_input_size, Some(31)); assert_eq!(model.input_processing_queries, 1); assert_eq!(model.output_generation_queries, 2); - assert_eq!(model.source_observed_at_unix_ms, 5); + assert_eq!(model.source_observed_at_unix_ms, 5_000); + assert!(model.complete); + } + } + + #[test] + fn counter_reset_skips_one_rate_sample_then_recovers() { + let identity = pool(1); + let mut translator = RelayLoadTranslator::default(); + translator + .translate(proto::LoadSnapshot { + metadata: Some(metadata()), + pools: vec![load_pool(identity.clone())], + models: vec![model_load("model-a", vec![identity.clone()])], + }) + .unwrap(); + + let mut reset = model_load("model-a", vec![identity.clone()]); + reset.input_tokens_total = Some(4); + reset.output_tokens_total = Some(2); + reset.source_observed_at_unix_ms = 6_000; + let reset = translator + .translate(proto::LoadSnapshot { + metadata: Some(metadata()), + pools: vec![load_pool(identity.clone())], + models: vec![reset], + }) + .unwrap(); + assert_eq!(reset.relay_models.get("model-a"), Some(&true)); + assert!(reset.stats.models.iter().all(|model| !model.complete)); + + let mut recovered = model_load("model-a", vec![identity.clone()]); + recovered.input_tokens_total = Some(14); + recovered.output_tokens_total = Some(7); + recovered.source_observed_at_unix_ms = 7_000; + let recovered = translator + .translate(proto::LoadSnapshot { + metadata: Some(metadata()), + pools: vec![load_pool(identity)], + models: vec![recovered], + }) + .unwrap(); + for model in recovered.stats.models { assert!(model.complete); + assert_eq!(model.input_tps, Some(10.0)); + assert_eq!(model.output_tps, 5.0); } } + #[test] + fn duplicate_load_model_rejects_even_an_incomplete_snapshot() { + let identity = pool(1); + let mut first = model_load("model-a", vec![identity.clone()]); + first.status = proto::DataStatus::Unavailable as i32; + first.output_tokens_total = None; + let mut second = first.clone(); + second.serving_pools.clear(); + + let result = RelayLoadTranslator::default().translate(proto::LoadSnapshot { + metadata: Some(metadata()), + pools: vec![load_pool(identity)], + models: vec![first, second], + }); + + assert!(result.is_err()); + } + #[test] fn unknown_exact_input_gauges_do_not_deactivate_the_model() { let identity = pool(1); + let mut translator = RelayLoadTranslator::default(); + let mut baseline = model_load("model-a", vec![identity.clone()]); + baseline.input_tokens_total = Some(0); + baseline.output_tokens_total = Some(0); + baseline.source_observed_at_unix_ms = 4_000; + translator + .translate(proto::LoadSnapshot { + metadata: Some(metadata()), + pools: vec![load_pool(identity.clone())], + models: vec![baseline], + }) + .unwrap(); let mut model = model_load("model-a", vec![identity.clone()]); model.pending_first_output_input_tokens = None; model.live_input_tokens = None; let snapshot = proto::LoadSnapshot { metadata: Some(metadata()), - window_ms: 1_000, pools: vec![load_pool(identity)], models: vec![model], }; - let translated = load_snapshot_from_proto(snapshot).unwrap(); + let translated = translator.translate(snapshot).unwrap(); assert_eq!(translated.relay_models.get("model-a"), Some(&true)); for stats in translated.stats.models { @@ -541,11 +683,10 @@ mod tests { let snapshot = proto::LoadSnapshot { metadata: Some(metadata()), - window_ms: 1_000, pools: vec![load_pool(identity)], models: vec![relay_only, model_load("frontend-only", Vec::new())], }; - let translated = load_snapshot_from_proto(snapshot).unwrap(); + let translated = RelayLoadTranslator::default().translate(snapshot).unwrap(); assert_eq!(translated.relay_models.get("relay-only"), Some(&false)); assert_eq!( @@ -582,10 +723,13 @@ mod tests { let load_snapshot = proto::LoadSnapshot { metadata: Some(metadata()), - window_ms: 1_000, pools: Vec::new(), models: vec![model_load("model-a", vec![pool(9)])], }; - assert!(load_snapshot_from_proto(load_snapshot).is_err()); + assert!( + RelayLoadTranslator::default() + .translate(load_snapshot) + .is_err() + ); } } diff --git a/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs b/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs index 28e0dd702..ef5729300 100644 --- a/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs +++ b/src/libraries/rust/stargate/crates/stargate/tests/suite/integration.rs @@ -627,10 +627,11 @@ impl stats_proto::kv_dc_relay_server::KvDcRelay for EngineStatsGrpc { let _ = self.state.connected_tx.send(true); let model = self.state.model.clone(); let stream = async_stream::stream! { + let mut sample = 0_u64; loop { + sample += 1; yield Ok(stats_proto::LoadSnapshot { metadata: Some(relay_metadata()), - window_ms: 1_000, pools: vec![stats_proto::PoolLoad { pool: Some(stats_pool_identity()), role: stats_proto::WorkerRole::Aggregated as i32, @@ -650,16 +651,16 @@ impl stats_proto::kv_dc_relay_server::KvDcRelay for EngineStatsGrpc { input_processing_requests: Some(1), output_generation_requests: Some(2), serving_pools: vec![stats_pool_identity()], - requests_started: 4, - requests_completed: 1, - requests_failed: 0, - requests_cancelled: 0, - input_tokens: Some(31), - output_tokens: 20, + requests_started_total: Some(sample.saturating_mul(4)), + requests_completed_total: Some(sample), + requests_failed_total: Some(0), + requests_cancelled_total: Some(0), + input_tokens_total: Some(sample.saturating_mul(4)), + output_tokens_total: Some(sample.saturating_mul(2)), status: stats_proto::DataStatus::Complete as i32, expected_frontends: 1, observed_frontends: 1, - source_observed_at_unix_ms: 1, + source_observed_at_unix_ms: sample.saturating_mul(100), }], }); tokio::time::sleep(Duration::from_millis(100)).await;