Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 31 additions & 3 deletions lib/llm/src/lora/filtered_router.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
};
Comment thread
ishandhanani marked this conversation as resolved.
tracker.record_worker(worker_id, None, worker_type);
}
}

#[async_trait]
Expand All @@ -163,7 +178,14 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<Annotated<LLMEngineOutpu
// the inner load-aware push router so base traffic on a LoRA-enabled deployment keeps the
// unmodified hot path (no avail/free scans, no set allocation, no LoadGuard).
let Some(lora_name) = lora_name else {
return self.inner.generate(request).await;
let ((tracker, worker_id), stream) = self
.inner
.select_and_dispatch(request, |request, worker_id| {
Ok((request.tracker.take(), worker_id))
})
.await?;
Self::record_worker(tracker.as_deref(), worker_id);
return Ok(stream);
};

self.load_estimator.increment_load(&lora_name);
Expand Down Expand Up @@ -229,10 +251,16 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<Annotated<LLMEngineOutpu
// race where `target` disappears mid-dispatch reselects another replica-set worker rather
// than escaping to an arbitrary worker outside the placement table.
let candidate_set: std::collections::HashSet<u64> = 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,
Expand Down
183 changes: 157 additions & 26 deletions lib/llm/src/session_affinity/push_router.rs
Comment thread
ishandhanani marked this conversation as resolved.
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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 {
Expand Down Expand Up @@ -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 {
Expand All @@ -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<Arc<RequestTracker>>, 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<AffinityTarget>,
Expand Down Expand Up @@ -118,7 +135,8 @@ impl SessionAffinityPushRouter {
where
F: FnOnce(&mut PreprocessedRequest, AffinityTarget) -> Result<M, Error>,
{
self.inner
let ((metadata, tracker, target), stream) = self
.inner
.select_and_dispatch_exact(
request,
pinned_target.map(|target| target.worker_id),
Expand All @@ -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<M, F>(
Expand Down Expand Up @@ -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;
}

Expand All @@ -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))
}
}

Expand All @@ -230,7 +250,20 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<LlmResponse>, 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| {
Comment thread
ishandhanani marked this conversation as resolved.
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 {
Expand All @@ -239,7 +272,19 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<LlmResponse>, 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();
Expand All @@ -251,7 +296,7 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<LlmResponse>, 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,
Expand All @@ -260,17 +305,17 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<LlmResponse>, 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);
}

Expand All @@ -293,19 +338,20 @@ impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<LlmResponse>, 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)
}
}

Expand All @@ -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;

Expand Down Expand Up @@ -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();
Expand Down Expand Up @@ -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();
Expand Down Expand Up @@ -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
Expand All @@ -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();
Expand Down
Loading
Loading