diff --git a/lib/llm/src/lora/filtered_router.rs b/lib/llm/src/lora/filtered_router.rs index a0aa025b84fc..f5b76a102861 100644 --- a/lib/llm/src/lora/filtered_router.rs +++ b/lib/llm/src/lora/filtered_router.rs @@ -28,6 +28,9 @@ use crate::lora::filter::LoraFilter; use crate::lora::load_estimator::LoadEstimator; use crate::preprocessor::PreprocessedRequest; use crate::protocols::common::llm_backend::LLMEngineOutput; +use crate::protocols::common::timing::{ + RequestPhase, RequestTracker, WORKER_TYPE_DECODE, WORKER_TYPE_PREFILL, +}; /// Decrements the [`LoadEstimator`] counter for a LoRA when dropped. struct LoadGuard { @@ -143,6 +146,18 @@ impl LoraFilteredRouter { } } } + + fn record_worker(tracker: Option<&RequestTracker>, worker_id: u64) { + let Some(tracker) = tracker else { + return; + }; + let worker_type = if tracker.phase() == RequestPhase::Prefill { + WORKER_TYPE_PREFILL + } else { + WORKER_TYPE_DECODE + }; + tracker.record_worker(worker_id, None, worker_type); + } } #[async_trait] @@ -163,7 +178,14 @@ impl AsyncEngine, ManyOut, ManyOut = candidates.iter().copied().collect(); - let response_stream = self + let ((tracker, worker_id), response_stream) = self .inner - .direct_within(request, target, Some(&candidate_set)) + .direct_within_prepared( + request, + target, + Some(&candidate_set), + |request, worker_id| Ok((request.tracker.take(), worker_id)), + ) .await?; + Self::record_worker(tracker.as_deref(), worker_id); let tracking = LoadTrackingStream { inner: response_stream, _guard: guard, diff --git a/lib/llm/src/session_affinity/push_router.rs b/lib/llm/src/session_affinity/push_router.rs index d5efd7e00f14..0a2844ee0d6d 100644 --- a/lib/llm/src/session_affinity/push_router.rs +++ b/lib/llm/src/session_affinity/push_router.rs @@ -1,7 +1,7 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -use std::time::Duration; +use std::{sync::Arc, time::Duration}; use dynamo_runtime::pipeline::{ AsyncEngine, AsyncEngineContext, AsyncEngineContextProvider, Error, ManyOut, PushRouter, @@ -15,7 +15,9 @@ use super::{ }; use crate::{ preprocessor::PreprocessedRequest, - protocols::common::timing::{RequestPhase, WORKER_TYPE_DECODE, WORKER_TYPE_PREFILL}, + protocols::common::timing::{ + RequestPhase, RequestTracker, WORKER_TYPE_DECODE, WORKER_TYPE_PREFILL, + }, }; pub struct SessionAffinityPushRouter { @@ -45,8 +47,8 @@ impl SessionAffinityPushRouter { .unwrap_or(RequestPhase::Aggregated) } - fn record_target(request: &PreprocessedRequest, target: AffinityTarget) { - let Some(tracker) = request.tracker.as_ref() else { + fn record_target(tracker: Option<&RequestTracker>, target: AffinityTarget) { + let Some(tracker) = tracker else { return; }; let worker_type = if tracker.phase() == RequestPhase::Prefill { @@ -57,6 +59,21 @@ impl SessionAffinityPushRouter { tracker.record_worker(target.worker_id, target.dp_rank, worker_type); } + fn prepare_resolved_target( + request: &mut PreprocessedRequest, + requested: AffinityTarget, + worker_id: u64, + ) -> (Option>, AffinityTarget) { + let dp_rank = requested + .dp_rank + .filter(|_| worker_id == requested.worker_id); + request.routing_mut().dp_rank = dp_rank; + ( + request.tracker.take(), + AffinityTarget { worker_id, dp_rank }, + ) + } + fn direct_target( &self, explicit: Option, @@ -118,7 +135,8 @@ impl SessionAffinityPushRouter { where F: FnOnce(&mut PreprocessedRequest, AffinityTarget) -> Result, { - self.inner + let ((metadata, tracker, target), stream) = self + .inner .select_and_dispatch_exact( request, pinned_target.map(|target| target.worker_id), @@ -128,10 +146,13 @@ impl SessionAffinityPushRouter { dp_rank: None, }); debug_assert_eq!(target.worker_id, worker_id); - prepare(request, target) + let metadata = prepare(request, target)?; + Ok((metadata, request.tracker.take(), target)) }, ) - .await + .await?; + Self::record_target(tracker.as_deref(), target); + Ok((metadata, stream)) } pub async fn select_and_dispatch_prefill( @@ -176,10 +197,7 @@ impl SessionAffinityPushRouter { .query_target(&session_id, explicit)? .or(explicit); return self - .select_and_dispatch_exact_target(request, selected, move |request, target| { - Self::record_target(request, target); - prepare(request, target) - }) + .select_and_dispatch_exact_target(request, selected, prepare) .await; } @@ -199,19 +217,21 @@ impl SessionAffinityPushRouter { worker_id, dp_rank: rank, }; - Self::record_target(request, target); - Ok((prepare(request, target)?, target)) + let metadata = prepare(request, target)?; + Ok((metadata, request.tracker.take(), target)) }, ) .await; - let ((metadata, target), stream) = match dispatch { + let ((metadata, tracker, target), stream) = match dispatch { Ok(result) => result, Err(error) => { operation.invalidate(); return Err(error); } }; - Ok((metadata, operation.into_stream(target, stream)?)) + let stream = operation.into_stream(target, stream)?; + Self::record_target(tracker.as_deref(), target); + Ok((metadata, stream)) } } @@ -230,7 +250,20 @@ impl AsyncEngine, ManyOut, Error> None }; if !self.direct && session_id.is_none() { - return self.inner.generate(request).await; + let ((tracker, target), stream) = self + .inner + .select_and_dispatch(request, |request, worker_id| { + Ok(( + request.tracker.take(), + AffinityTarget { + worker_id, + dp_rank: None, + }, + )) + }) + .await?; + Self::record_target(tracker.as_deref(), target); + return Ok(stream); } let explicit = self.direct_target(explicit_target(&request, phase)?, phase)?; let Some(session_id) = session_id else { @@ -239,7 +272,19 @@ impl AsyncEngine, ManyOut, Error> "Direct routing requires an explicit {phase} target" ))); }; - return self.inner.direct(request, target.worker_id).await; + let ((tracker, target), stream) = self + .inner + .direct_within_prepared( + request, + target.worker_id, + None, + move |request, worker_id| { + Ok(Self::prepare_resolved_target(request, target, worker_id)) + }, + ) + .await?; + Self::record_target(tracker.as_deref(), target); + return Ok(stream); }; let is_query_only = request.get_annotation_value("query_instance_id").is_some(); @@ -251,7 +296,7 @@ impl AsyncEngine, ManyOut, Error> .query_target(&session_id, explicit)? .or(explicit); let rank = target.and_then(|target| target.dp_rank); - let (_, stream) = self + let ((tracker, target), stream) = self .inner .select_and_dispatch_exact( request, @@ -260,17 +305,17 @@ impl AsyncEngine, ManyOut, Error> if rank.is_some() { request.routing_mut().dp_rank = rank; } - Self::record_target( - request, + Ok(( + request.tracker.take(), AffinityTarget { worker_id, dp_rank: rank, }, - ); - Ok(()) + )) }, ) .await?; + Self::record_target(tracker.as_deref(), target); return Ok(stream); } @@ -293,19 +338,20 @@ impl AsyncEngine, ManyOut, Error> worker_id, dp_rank: rank, }; - Self::record_target(request, target); - Ok(target) + Ok((request.tracker.take(), target)) }, ) .await; - let (target, stream) = match dispatch { + let ((tracker, target), stream) = match dispatch { Ok(result) => result, Err(error) => { operation.invalidate(); return Err(error); } }; - operation.into_stream(target, stream) + let stream = operation.into_stream(target, stream)?; + Self::record_target(tracker.as_deref(), target); + Ok(stream) } } @@ -321,6 +367,7 @@ mod tests { use crate::protocols::common::{ extensions::{SESSION_AFFINITY_CONTEXT_KEY, SessionAffinityId}, preprocessor::RoutingHints, + timing::RequestTracker, }; use crate::session_affinity::AffinityAcquire; @@ -360,6 +407,30 @@ mod tests { .expect("test router must enable affinity") } + #[test] + fn direct_fallback_clears_stale_dp_rank() { + let tracker = Arc::new(RequestTracker::new()); + let mut content = request(Some(7), false); + content.routing_mut().dp_rank = Some(3); + content.tracker = Some(tracker.clone()); + + let (prepared_tracker, target) = SessionAffinityPushRouter::prepare_resolved_target( + &mut content, + AffinityTarget { + worker_id: 7, + dp_rank: Some(3), + }, + 8, + ); + + assert_eq!(content.routing.unwrap().dp_rank, None); + assert_eq!(tracker.prefill_worker_id(), None); + assert_eq!(tracker.decode_worker_id(), None); + SessionAffinityPushRouter::record_target(prepared_tracker.as_deref(), target); + assert_eq!(tracker.prefill_worker_id(), Some(8)); + assert_eq!(tracker.decode_worker_id(), Some(8)); + } + #[tokio::test] async fn session_affinity_disabled_simple_router_has_no_coordinator() { let runtime = Runtime::from_current().unwrap(); @@ -387,6 +458,62 @@ mod tests { runtime.shutdown(); } + #[tokio::test] + async fn failed_non_kv_dispatch_does_not_record_selected_worker() { + let runtime = Runtime::from_current().unwrap(); + let distributed = + DistributedRuntime::new(runtime.clone(), DistributedConfig::process_local()) + .await + .unwrap(); + let component = distributed + .namespace("session_affinity_worker_disclosure".to_string()) + .unwrap() + .component("workers".to_string()) + .unwrap(); + + for (index, mode) in [ + RouterMode::Random, + RouterMode::RoundRobin, + RouterMode::PowerOfTwoChoices, + RouterMode::LeastLoaded, + RouterMode::DeviceAwareWeighted, + RouterMode::Direct, + ] + .into_iter() + .enumerate() + { + let endpoint = component.endpoint(format!("mode-{index}")); + let client = endpoint.client().await.unwrap(); + endpoint.register_endpoint_instance().await.unwrap(); + let worker_id = client.wait_for_instances().await.unwrap()[0].id(); + let inner = PushRouter::from_client(client, mode).await.unwrap(); + let router = + SessionAffinityPushRouter::new(inner, None, mode.is_direct_routing()).unwrap(); + let tracker = Arc::new(RequestTracker::new()); + let mut content = request(mode.is_direct_routing().then_some(worker_id), false); + content.tracker = Some(tracker.clone()); + + let _ = tokio::time::timeout( + Duration::from_millis(100), + router.generate(Context::new(content)), + ) + .await; + + assert_eq!( + tracker.prefill_worker_id(), + None, + "{mode:?} must not disclose a worker before dispatch succeeds" + ); + assert_eq!( + tracker.decode_worker_id(), + None, + "{mode:?} must not disclose a worker before dispatch succeeds" + ); + } + + runtime.shutdown(); + } + #[tokio::test] async fn session_affinity_simple_modes_rollback_failed_initialization() { let runtime = Runtime::from_current().unwrap(); @@ -536,6 +663,8 @@ mod tests { let mut content = request(None, false); content.routing_mut().prefill_worker_id = Some(worker_id); content.routing_mut().prefill_dp_rank = Some(0); + let tracker = Arc::new(RequestTracker::new()); + content.tracker = Some(tracker.clone()); let mut observed = None; let error = router @@ -548,6 +677,8 @@ mod tests { assert!(error.to_string().contains("stop before dispatch")); assert_eq!(observed, Some(expected)); + assert_eq!(tracker.prefill_worker_id(), None); + assert_eq!(tracker.decode_worker_id(), None); } runtime.shutdown(); diff --git a/lib/runtime/src/pipeline/network/egress/push_router.rs b/lib/runtime/src/pipeline/network/egress/push_router.rs index 217739cc5c84..c24130765c22 100644 --- a/lib/runtime/src/pipeline/network/egress/push_router.rs +++ b/lib/runtime/src/pipeline/network/egress/push_router.rs @@ -79,6 +79,15 @@ impl OccupancyPermit { } } + fn retarget(&mut self, instance_id: u64) { + if self.instance_id == instance_id { + return; + } + self.state.increment(instance_id); + self.state.decrement(self.instance_id); + self.instance_id = instance_id; + } + fn into_tracked_stream(mut self, stream: ManyOut) -> ManyOut { self.armed = false; let engine_ctx = stream.context(); @@ -668,6 +677,19 @@ where /// Issue a request to the next available instance in a round-robin fashion pub async fn round_robin(&self, request: SingleIn) -> anyhow::Result> { + self.round_robin_prepared(request, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + async fn round_robin_prepared( + &self, + request: SingleIn, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { let counter = self.round_robin_counter.fetch_add(1, Ordering::Relaxed) as usize; let (instance_id, candidate_count) = { @@ -685,12 +707,25 @@ where "Selected worker" ); - self.generate_with_fault_detection(instance_id, request, TransportFallback::Allow) + self.dispatch_selected(instance_id, request, None, prepare) .await } /// Issue a request to a random endpoint pub async fn random(&self, request: SingleIn) -> anyhow::Result> { + self.random_prepared(request, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + async fn random_prepared( + &self, + request: SingleIn, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { let (instance_id, candidate_count) = { let routing_instances = self.client.routing_instances(); let count = routing_instances.free_ids().len(); @@ -707,13 +742,26 @@ where "Selected worker" ); - self.generate_with_fault_detection(instance_id, request, TransportFallback::Allow) + self.dispatch_selected(instance_id, request, None, prepare) .await } /// Issue a request using power-of-two-choices: pick 2 random healthy workers, /// route to the one with fewer in-flight requests. pub async fn power_of_two_choices(&self, request: SingleIn) -> anyhow::Result> { + self.power_of_two_choices_prepared(request, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + async fn power_of_two_choices_prepared( + &self, + request: SingleIn, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { let state = self.occupancy_state()?; let instance_id = { let routing_instances = self.client.routing_instances(); @@ -724,14 +772,8 @@ where }; state.increment(instance_id); let permit = OccupancyPermit::new(state, instance_id); - - match self - .generate_with_fault_detection(instance_id, request, TransportFallback::Allow) + self.dispatch_selected(instance_id, request, Some(permit), prepare) .await - { - Ok(stream) => Ok(permit.into_tracked_stream(stream)), - Err(err) => Err(err), - } } /// Issue a request to a specific endpoint @@ -753,6 +795,23 @@ where instance_id: u64, allowed_fallback: Option<&HashSet>, ) -> anyhow::Result> { + self.direct_within_prepared(request, instance_id, allowed_fallback, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + /// Like [`Self::direct_within`], but prepares the request after transport resolution and + /// returns the preparation metadata alongside the response stream. + pub async fn direct_within_prepared( + &self, + request: SingleIn, + instance_id: u64, + allowed_fallback: Option<&HashSet>, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { // When fault detection is disabled, check the raw discovery list // (not filtered by report_instance_down) so transient failures // don't poison the instance for subsequent retries. @@ -785,7 +844,7 @@ where let fallback = allowed_fallback .map(TransportFallback::Within) .unwrap_or(TransportFallback::Allow); - self.generate_with_fault_detection(instance_id, request, fallback) + self.generate_with_fault_detection_prepared(instance_id, request, fallback, prepare) .await } @@ -825,6 +884,63 @@ where Ok((metadata, stream)) } + /// Select a worker using the configured routing mode, prepare the request with the worker + /// that survives transport resolution, then dispatch with normal fallback behavior. + pub async fn select_and_dispatch( + &self, + request: SingleIn, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { + match self.router_mode { + RouterMode::Random => self.random_prepared(request, prepare).await, + RouterMode::RoundRobin => self.round_robin_prepared(request, prepare).await, + RouterMode::PowerOfTwoChoices => { + self.power_of_two_choices_prepared(request, prepare).await + } + RouterMode::LeastLoaded => self.least_loaded_prepared(request, prepare).await, + RouterMode::DeviceAwareWeighted => { + self.device_aware_weighted_prepared(request, prepare).await + } + RouterMode::KV => anyhow::bail!("KV routing should not call select_and_dispatch"), + RouterMode::Direct => anyhow::bail!( + "Direct routing should use direct_within_prepared instead of select_and_dispatch" + ), + } + } + + async fn dispatch_selected( + &self, + instance_id: u64, + request: SingleIn, + mut permit: Option, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { + let (metadata, stream) = self + .generate_with_fault_detection_prepared( + instance_id, + request, + TransportFallback::Allow, + |request, resolved_instance_id| { + if let Some(permit) = permit.as_mut() { + permit.retarget(resolved_instance_id); + } + prepare(request, resolved_instance_id) + }, + ) + .await?; + let stream = match permit { + Some(permit) => permit.into_tracked_stream(stream), + None => stream, + }; + Ok((metadata, stream)) + } + /// Issue a request using device-aware weighted routing. /// /// Instances are partitioned by device type (CPU vs non-CPU), then the router @@ -834,6 +950,19 @@ where /// If only one device class exists (all CPU or all non-CPU), this naturally /// degenerates to least-loaded routing over the available instances. pub async fn device_aware_weighted(&self, request: SingleIn) -> anyhow::Result> { + self.device_aware_weighted_prepared(request, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + async fn device_aware_weighted_prepared( + &self, + request: SingleIn, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { let state = self.occupancy_state()?; let routing_instances = self.client.routing_instances(); let instance_ids = routing_instances.free_ids().to_vec(); @@ -879,16 +1008,8 @@ where "Selected worker" ); - match self - .generate_with_fault_detection(instance_id, request, TransportFallback::Allow) + self.dispatch_selected(instance_id, request, permit, prepare) .await - { - Ok(stream) => Ok(match permit { - Some(permit) => permit.into_tracked_stream(stream), - None => stream, - }), - Err(err) => Err(err), - } } fn device_aware_candidates( @@ -956,6 +1077,19 @@ where /// Issue a request to the instance with the fewest active connections. pub async fn least_loaded(&self, request: SingleIn) -> anyhow::Result> { + self.least_loaded_prepared(request, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + async fn least_loaded_prepared( + &self, + request: SingleIn, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { let state = self.occupancy_state()?; let routing_instances = self.client.routing_instances(); let instance_ids = routing_instances.free_ids().to_vec(); @@ -972,13 +1106,8 @@ where "Selected worker" ); - match self - .generate_with_fault_detection(instance_id, request, TransportFallback::Allow) + self.dispatch_selected(instance_id, request, Some(permit), prepare) .await - { - Ok(stream) => Ok(permit.into_tracked_stream(stream)), - Err(err) => Err(err), - } } /// Select the next worker according to the routing mode. @@ -1185,6 +1314,21 @@ where request: SingleIn, fallback: TransportFallback<'_>, ) -> anyhow::Result> { + self.generate_with_fault_detection_prepared(instance_id, request, fallback, |_, _| Ok(())) + .await + .map(|(_, stream)| stream) + } + + async fn generate_with_fault_detection_prepared( + &self, + instance_id: u64, + mut request: SingleIn, + fallback: TransportFallback<'_>, + prepare: F, + ) -> anyhow::Result<(M, ManyOut)> + where + F: FnOnce(&mut T, u64) -> anyhow::Result, + { let route_start = Instant::now(); let request_id = request.id().to_string(); let route_span = if matches!(self.router_mode, RouterMode::KV) { @@ -1202,6 +1346,7 @@ where let (instance_id, address, transport_kind, instance) = self.resolve_transport(instance_id, fallback)?; + let metadata = prepare(&mut request, instance_id)?; let request = request.map(|req| AddressedRequest::with_instance(req, address, instance)); STAGE_DURATION_SECONDS @@ -1214,7 +1359,8 @@ where .generate(request) .instrument(route_span) .await; - self.wrap_with_fault_detection(stream, instance_id) + let stream = self.wrap_with_fault_detection(stream, instance_id)?; + Ok((metadata, stream)) } /// Reject early if the selected worker is overloaded and fault detection @@ -2403,6 +2549,49 @@ mod tests { rt.shutdown(); } + #[tokio::test] + async fn prepared_dispatch_observes_worker_after_transport_fallback() { + let rt = Runtime::from_current().unwrap(); + let drt = DistributedRuntime::new(rt.clone(), DistributedConfig::process_local()) + .await + .unwrap(); + let endpoint = drt + .namespace("test_prepared_transport_fallback".to_string()) + .unwrap() + .component("test_component".to_string()) + .unwrap() + .endpoint("test_endpoint".to_string()); + let client = endpoint.client().await.unwrap(); + endpoint.register_endpoint_instance().await.unwrap(); + let real_id = client.wait_for_instances().await.unwrap()[0].id(); + let stale_id = real_id.wrapping_add(1); + client.override_instance_avail(vec![stale_id, real_id]); + let router = PushRouter::::from_client(client, RouterMode::LeastLoaded) + .await + .unwrap(); + let state = router.occupancy_state.clone().unwrap(); + state.increment(real_id); + let state_for_prepare = state.clone(); + let observed = Arc::new(AtomicU64::new(0)); + let observed_for_prepare = observed.clone(); + + let _ = tokio::time::timeout( + std::time::Duration::from_millis(100), + router.select_and_dispatch(SingleIn::new(42), move |_, worker_id| { + assert_eq!(state_for_prepare.load(stale_id), 0); + assert_eq!(state_for_prepare.load(worker_id), 2); + observed_for_prepare.store(worker_id, Ordering::Relaxed); + Ok(()) + }), + ) + .await; + + assert_eq!(observed.load(Ordering::Relaxed), real_id); + assert_eq!(state.load(real_id), 1); + state.decrement(real_id); + rt.shutdown(); + } + /// When no instances are available at all (both primary and fallback), /// the router should return a clear error. #[tokio::test]