use std::{sync::Arc, time::Duration};
use dynamo_runtime::pipeline::{
AsyncEngine, AsyncEngineContext, AsyncEngineContextProvider, Error, ManyOut, PushRouter,
SingleIn, async_trait as pipeline_async_trait,
};
use super::{
AffinityCoordinator, AffinityTarget, LlmResponse,
coordinator::{affinity_id, invalid_argument},
explicit_target,
};
use crate::{
preprocessor::PreprocessedRequest,
protocols::common::timing::{
RequestPhase, RequestTracker, WORKER_TYPE_DECODE, WORKER_TYPE_PREFILL,
},
};
pub struct SessionAffinityPushRouter {
inner: PushRouter<PreprocessedRequest, LlmResponse>,
affinity: Option<AffinityCoordinator>,
direct: bool,
}
impl SessionAffinityPushRouter {
pub fn new(
inner: PushRouter<PreprocessedRequest, LlmResponse>,
ttl: Option<Duration>,
direct: bool,
) -> Result<Self, Error> {
Ok(Self {
inner,
affinity: ttl.map(AffinityCoordinator::new).transpose()?,
direct,
})
}
fn phase(request: &PreprocessedRequest) -> RequestPhase {
request
.tracker
.as_ref()
.map(|tracker| tracker.phase())
.unwrap_or(RequestPhase::Aggregated)
}
fn record_target(tracker: Option<&RequestTracker>, target: AffinityTarget) {
let Some(tracker) = tracker else {
return;
};
let worker_type = if tracker.phase() == RequestPhase::Prefill {
WORKER_TYPE_PREFILL
} else {
WORKER_TYPE_DECODE
};
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>,
phase: RequestPhase,
) -> Result<Option<AffinityTarget>, Error> {
if !self.direct {
return Ok(explicit);
}
explicit.map(Some).ok_or_else(|| {
invalid_argument(format!(
"worker ID required for {phase} request in Direct routing mode"
))
})
}
pub fn peek_next_worker(&self) -> Option<u64> {
self.inner.peek_next_worker()
}
async fn acquire_routable(
&self,
session_id: &crate::protocols::common::extensions::SessionAffinityId,
explicit: Option<AffinityTarget>,
request_context: &dyn AsyncEngineContext,
) -> Result<super::AffinityAcquire, Error> {
let affinity = self
.affinity
.as_ref()
.expect("affinity acquisition requires an enabled coordinator");
let operation = affinity
.acquire_with_context(session_id, explicit, request_context)
.await?;
let Some(target) = operation.target() else {
return Ok(operation);
};
if self
.inner
.client
.instance_ids_avail()
.contains(&target.worker_id)
{
return Ok(operation);
}
operation.invalidate();
affinity
.acquire_with_context(session_id, explicit, request_context)
.await
}
async fn select_and_dispatch_prefill_exact<M, F>(
&self,
request: SingleIn<PreprocessedRequest>,
pinned_worker: Option<u64>,
dp_rank: Option<u32>,
prepare: F,
) -> Result<(M, ManyOut<LlmResponse>), Error>
where
F: FnOnce(&mut PreprocessedRequest, u64, Option<u32>) -> Result<M, Error>,
{
let ((metadata, tracker, target), stream) = self
.inner
.select_and_dispatch_exact(request, pinned_worker, move |request, worker_id| {
let target = AffinityTarget { worker_id, dp_rank };
let metadata = prepare(request, worker_id, dp_rank)?;
Ok((metadata, request.tracker.take(), target))
})
.await?;
Self::record_target(tracker.as_deref(), target);
Ok((metadata, stream))
}
pub async fn select_and_dispatch_prefill<M, F>(
&self,
request: SingleIn<PreprocessedRequest>,
prepare: F,
) -> Result<(M, ManyOut<LlmResponse>), Error>
where
F: FnOnce(&mut PreprocessedRequest, u64, Option<u32>) -> Result<M, Error>,
{
let session_id = if self.affinity.is_some() {
affinity_id(&request)?
} else {
None
};
if !self.direct && session_id.is_none() {
let pinned_worker = phase_worker_id(&request, RequestPhase::Prefill);
return self
.select_and_dispatch_prefill_exact(request, pinned_worker, None, prepare)
.await;
}
let explicit = self.direct_target(
explicit_target(&request, RequestPhase::Prefill)?,
RequestPhase::Prefill,
)?;
let Some(session_id) = session_id else {
let Some(pinned_worker) = explicit else {
return Err(invalid_argument(
"Direct routing requires an explicit prefill target",
));
};
return self
.select_and_dispatch_prefill_exact(
request,
Some(pinned_worker.worker_id),
None,
prepare,
)
.await;
};
let is_query_only = request.get_annotation_value("query_instance_id").is_some();
if is_query_only {
let selected = self
.affinity
.as_ref()
.expect("affinity query requires an enabled coordinator")
.query_target(&session_id, explicit)?
.or(explicit);
let rank = selected.and_then(|target| target.dp_rank);
return self
.select_and_dispatch_prefill_exact(
request,
selected.map(|target| target.worker_id),
rank,
prepare,
)
.await;
}
let request_context = request.context();
let operation = self
.acquire_routable(&session_id, explicit, request_context.as_ref())
.await?;
let selected = operation.target().or(explicit);
let rank = selected.and_then(|target| target.dp_rank);
let dispatch = self
.inner
.select_and_dispatch_exact(
request,
selected.map(|target| target.worker_id),
move |request, worker_id| {
let target = AffinityTarget {
worker_id,
dp_rank: rank,
};
let metadata = prepare(request, worker_id, rank)?;
Ok((metadata, request.tracker.take(), target))
},
)
.await;
let ((metadata, tracker, target), stream) = match dispatch {
Ok(result) => result,
Err(error) => {
operation.invalidate();
return Err(error);
}
};
let stream = operation.into_stream(target, stream)?;
Self::record_target(tracker.as_deref(), target);
Ok((metadata, stream))
}
}
#[pipeline_async_trait]
impl AsyncEngine<SingleIn<PreprocessedRequest>, ManyOut<LlmResponse>, Error>
for SessionAffinityPushRouter
{
async fn generate(
&self,
request: SingleIn<PreprocessedRequest>,
) -> Result<ManyOut<LlmResponse>, Error> {
let phase = Self::phase(&request);
let session_id = if self.affinity.is_some() {
affinity_id(&request)?
} else {
None
};
if !self.direct && session_id.is_none() {
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 {
let Some(target) = explicit else {
return Err(invalid_argument(format!(
"Direct routing requires an explicit {phase} target"
)));
};
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();
if is_query_only {
let target = self
.affinity
.as_ref()
.expect("affinity query requires an enabled coordinator")
.query_target(&session_id, explicit)?
.or(explicit);
let rank = target.and_then(|target| target.dp_rank);
let ((tracker, target), stream) = self
.inner
.select_and_dispatch_exact(
request,
target.map(|target| target.worker_id),
move |request, worker_id| {
if rank.is_some() {
request.routing_mut().dp_rank = rank;
}
Ok((
request.tracker.take(),
AffinityTarget {
worker_id,
dp_rank: rank,
},
))
},
)
.await?;
Self::record_target(tracker.as_deref(), target);
return Ok(stream);
}
let request_context = request.context();
let operation = self
.acquire_routable(&session_id, explicit, request_context.as_ref())
.await?;
let selected = operation.target().or(explicit);
let rank = selected.and_then(|target| target.dp_rank);
let dispatch = self
.inner
.select_and_dispatch_exact(
request,
selected.map(|target| target.worker_id),
move |request, worker_id| {
if rank.is_some() {
request.routing_mut().dp_rank = rank;
}
let target = AffinityTarget {
worker_id,
dp_rank: rank,
};
Ok((request.tracker.take(), target))
},
)
.await;
let ((tracker, target), stream) = match dispatch {
Ok(result) => result,
Err(error) => {
operation.invalidate();
return Err(error);
}
};
let stream = operation.into_stream(target, stream)?;
Self::record_target(tracker.as_deref(), target);
Ok(stream)
}
}
fn phase_worker_id(request: &PreprocessedRequest, phase: RequestPhase) -> Option<u64> {
let routing = request.routing.as_ref()?;
match phase {
RequestPhase::Prefill => routing.prefill_worker_id.or(routing.backend_instance_id),
RequestPhase::Decode => routing.decode_worker_id.or(routing.backend_instance_id),
RequestPhase::Aggregated => routing.decode_worker_id.or(routing.backend_instance_id),
}
}
#[cfg(test)]
mod tests {
use dynamo_runtime::{
DistributedRuntime, Runtime,
distributed::DistributedConfig,
pipeline::{Context, RouterMode},
};
use super::*;
use crate::protocols::common::{
extensions::{SESSION_AFFINITY_CONTEXT_KEY, SessionAffinityId},
preprocessor::RoutingHints,
timing::RequestTracker,
};
use crate::session_affinity::AffinityAcquire;
fn request(worker_id: Option<u64>, query_only: bool) -> PreprocessedRequest {
PreprocessedRequest::builder()
.model("test".to_string())
.token_ids(vec![1, 2, 3])
.stop_conditions(Default::default())
.sampling_options(Default::default())
.output_options(Default::default())
.annotations(if query_only {
vec!["query_instance_id:true".to_string()]
} else {
Vec::new()
})
.routing(worker_id.map(|worker_id| RoutingHints {
backend_instance_id: Some(worker_id),
..Default::default()
}))
.build()
.unwrap()
}
fn affinity_request(worker_id: Option<u64>, query_only: bool) -> SingleIn<PreprocessedRequest> {
let mut request = Context::new(request(worker_id, query_only));
request.insert(
SESSION_AFFINITY_CONTEXT_KEY,
SessionAffinityId::new("adapter-session"),
);
request
}
fn affinity(router: &SessionAffinityPushRouter) -> &AffinityCoordinator {
router
.affinity
.as_ref()
.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();
let distributed =
DistributedRuntime::new(runtime.clone(), DistributedConfig::process_local())
.await
.unwrap();
let client = distributed
.namespace("session_affinity_disabled".to_string())
.unwrap()
.component("workers".to_string())
.unwrap()
.endpoint("generate")
.client()
.await
.unwrap();
let inner = PushRouter::from_client(client, RouterMode::RoundRobin)
.await
.unwrap();
let router = SessionAffinityPushRouter::new(inner, None, false).unwrap();
assert!(router.affinity.is_none());
drop(router);
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();
let distributed =
DistributedRuntime::new(runtime.clone(), DistributedConfig::process_local())
.await
.unwrap();
let namespace = distributed
.namespace("session_affinity_adapters".to_string())
.unwrap();
let component = namespace.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();
let inner = PushRouter::from_client(client, mode).await.unwrap();
let router = SessionAffinityPushRouter::new(
inner,
Some(Duration::from_secs(10)),
mode.is_direct_routing(),
)
.unwrap();
let worker_id = mode.is_direct_routing().then_some(99);
assert!(
router
.generate(affinity_request(worker_id, false))
.await
.is_err()
);
assert_eq!(
affinity(&router).entry_count(),
0,
"failed {mode:?} dispatch must release initialization"
);
}
runtime.shutdown();
}
#[tokio::test]
async fn session_affinity_query_and_direct_validation_do_not_create_state() {
let runtime = Runtime::from_current().unwrap();
let distributed =
DistributedRuntime::new(runtime.clone(), DistributedConfig::process_local())
.await
.unwrap();
let namespace = distributed
.namespace("session_affinity_read_only".to_string())
.unwrap();
let component = namespace.component("workers".to_string()).unwrap();
let client = component
.endpoint("query".to_string())
.client()
.await
.unwrap();
let inner = PushRouter::from_client(client, RouterMode::RoundRobin)
.await
.unwrap();
let router =
SessionAffinityPushRouter::new(inner, Some(Duration::from_secs(10)), false).unwrap();
assert!(router.generate(affinity_request(None, true)).await.is_err());
assert_eq!(affinity(&router).entry_count(), 0);
assert!(
router
.select_and_dispatch_prefill(affinity_request(None, true), |_, _, _| Ok(()))
.await
.is_err()
);
assert_eq!(affinity(&router).entry_count(), 0);
let client = component
.endpoint("direct".to_string())
.client()
.await
.unwrap();
let inner = PushRouter::from_client(client, RouterMode::Direct)
.await
.unwrap();
let router =
SessionAffinityPushRouter::new(inner, Some(Duration::from_secs(10)), true).unwrap();
let error = router
.generate(affinity_request(None, false))
.await
.unwrap_err();
assert!(error.to_string().contains("worker ID required"));
assert_eq!(affinity(&router).entry_count(), 0);
let error = router
.generate(Context::new(request(None, false)))
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("worker ID required for aggregated request in Direct routing mode")
);
let error = router
.select_and_dispatch_prefill(Context::new(request(None, false)), |_, _, _| Ok(()))
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("worker ID required for prefill request in Direct routing mode")
);
assert_eq!(affinity(&router).entry_count(), 0);
let mut decode_only = request(None, false);
decode_only.routing_mut().decode_worker_id = Some(99);
assert_eq!(
phase_worker_id(&decode_only, RequestPhase::Aggregated),
Some(99)
);
runtime.shutdown();
}
#[tokio::test]
async fn failed_prefill_preparation_does_not_record_selected_worker() {
let runtime = Runtime::from_current().unwrap();
let distributed =
DistributedRuntime::new(runtime.clone(), DistributedConfig::process_local())
.await
.unwrap();
let endpoint = distributed
.namespace("session_affinity_prefill_target".to_string())
.unwrap()
.component("workers".to_string())
.unwrap()
.endpoint("prefill".to_string());
let client = endpoint.client().await.unwrap();
endpoint.register_endpoint_instance().await.unwrap();
let worker_id = client.wait_for_instances().await.unwrap()[0].id();
for (mode, direct) in [(RouterMode::Direct, true), (RouterMode::RoundRobin, false)] {
let inner = PushRouter::from_client(client.clone(), mode).await.unwrap();
let router = SessionAffinityPushRouter::new(inner, None, direct).unwrap();
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
.select_and_dispatch_prefill(Context::new(content), |_, worker_id, dp_rank| {
observed = Some((worker_id, dp_rank));
Err::<(), _>(anyhow::anyhow!("stop before dispatch"))
})
.await
.unwrap_err();
assert!(error.to_string().contains("stop before dispatch"));
assert_eq!(observed, Some((worker_id, None)));
assert_eq!(tracker.prefill_worker_id(), None);
assert_eq!(tracker.decode_worker_id(), None);
}
runtime.shutdown();
}
#[tokio::test]
async fn session_affinity_unavailable_target_is_invalidated() {
let runtime = Runtime::from_current().unwrap();
let distributed =
DistributedRuntime::new(runtime.clone(), DistributedConfig::process_local())
.await
.unwrap();
let namespace = distributed
.namespace("session_affinity_unavailable".to_string())
.unwrap();
let endpoint = namespace
.component("workers".to_string())
.unwrap()
.endpoint("generate".to_string());
let client = endpoint.client().await.unwrap();
let inner = PushRouter::from_client(client, RouterMode::RoundRobin)
.await
.unwrap();
let router =
SessionAffinityPushRouter::new(inner, Some(Duration::from_secs(10)), false).unwrap();
let session_id = SessionAffinityId::new("adapter-session");
let AffinityAcquire::Initialize(initializer) =
affinity(&router).acquire(&session_id, None).await.unwrap()
else {
panic!("first request must initialize");
};
drop(
initializer
.commit(AffinityTarget {
worker_id: 99,
dp_rank: None,
})
.unwrap(),
);
assert!(
router
.generate(affinity_request(None, false))
.await
.is_err()
);
assert_eq!(
affinity(&router).query_target(&session_id, None).unwrap(),
None
);
runtime.shutdown();
}
}