diff --git a/lib/llm/src/lib.rs b/lib/llm/src/lib.rs index 4e5281ed513d..d902aaf60b9d 100644 --- a/lib/llm/src/lib.rs +++ b/lib/llm/src/lib.rs @@ -34,6 +34,7 @@ pub mod recorder; pub mod request_template; pub mod request_trace; pub mod session_affinity; +pub(crate) mod session_placement; pub mod telemetry; pub use dynamo_tokenizers as tokenizers; pub use dynamo_tokenizers::{file_json_field, log_json_err}; diff --git a/lib/llm/src/session_affinity/coordinator.rs b/lib/llm/src/session_affinity/coordinator.rs index cc3a5ebd25d2..5e47f0a378b3 100644 --- a/lib/llm/src/session_affinity/coordinator.rs +++ b/lib/llm/src/session_affinity/coordinator.rs @@ -3,24 +3,11 @@ use std::{ pin::Pin, - sync::{ - Arc, Weak, - atomic::{AtomicU64, AtomicUsize, Ordering}, - }, + sync::Arc, task::{Context, Poll}, time::Duration, }; -use dashmap::{DashMap, mapref::entry::Entry}; -use dynamo_runtime::{ - engine::{AsyncEngineContext, AsyncEngineContextProvider}, - error::{DynamoError, ErrorType}, - pipeline::{Error, ManyOut, ResponseStream}, -}; -use futures::Stream; -use tokio::{sync::Notify, time::Instant}; -use tokio_util::sync::CancellationToken; - use super::{ LlmResponse, MAX_SESSION_AFFINITY_ENTRIES, MAX_SESSION_AFFINITY_ID_BYTES, MAX_SESSION_AFFINITY_TTL_SECS, @@ -31,7 +18,17 @@ use crate::{ extensions::{SESSION_AFFINITY_CONTEXT_KEY, SessionAffinityId}, timing::RequestPhase, }, + session_placement::{ + PlacementAcquire, PlacementInitialization, PlacementLease, SessionPlacement, + SessionPlacementConfig, SessionPlacementError, TargetGeneration, + }, }; +use dynamo_runtime::{ + engine::{AsyncEngineContext, AsyncEngineContextProvider}, + error::{DynamoError, ErrorType}, + pipeline::{Error, ManyOut, ResponseStream}, +}; +use futures::Stream; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct AffinityTarget { @@ -39,42 +36,10 @@ pub struct AffinityTarget { pub dp_rank: Option, } -enum AffinityEntry { - Initializing { - revision: u64, - notify: Arc, - }, - Bound { - target: AffinityTarget, - revision: u64, - active_leases: usize, - idle_deadline: Instant, - }, -} - -struct AffinityCoordinatorInner { - entries: DashMap, - ttl: Duration, - max_entries: usize, - max_session_id_bytes: usize, - entry_count: AtomicUsize, - next_revision: AtomicU64, - cancel: CancellationToken, - #[cfg(test)] - reaper_started: Arc, - #[cfg(test)] - waiter_observed: Arc, -} - -impl Drop for AffinityCoordinatorInner { - fn drop(&mut self) { - self.cancel.cancel(); - } -} - #[derive(Clone)] pub struct AffinityCoordinator { - inner: Arc, + placement: SessionPlacement, + max_session_id_bytes: usize, } impl AffinityCoordinator { @@ -98,57 +63,18 @@ impl AffinityCoordinator { "session affinity TTL must be between 1 and {MAX_SESSION_AFFINITY_TTL_SECS} seconds" ))); } - let inner = Arc::new(AffinityCoordinatorInner { - entries: DashMap::new(), - ttl, + let placement = SessionPlacement::new(SessionPlacementConfig { + idle_ttl: ttl, + // Existing worker affinity keeps initialization until its owner finishes or drops. + initialization_timeout: None, max_entries, + max_key_bytes: max_session_id_bytes, + }) + .map_err(map_placement_error)?; + Ok(Self { + placement, max_session_id_bytes, - entry_count: AtomicUsize::new(0), - next_revision: AtomicU64::new(1), - cancel: CancellationToken::new(), - #[cfg(test)] - reaper_started: Arc::new(Notify::new()), - #[cfg(test)] - waiter_observed: Arc::new(Notify::new()), - }); - Self::spawn_reaper(&inner); - Ok(Self { inner }) - } - - fn spawn_reaper(inner: &Arc) { - let weak = Arc::downgrade(inner); - let cancel = inner.cancel.clone(); - let period = inner.ttl.min(Duration::from_secs(30)); - #[cfg(test)] - let reaper_started = inner.reaper_started.clone(); - tokio::spawn(async move { - #[cfg(test)] - reaper_started.notify_one(); - loop { - tokio::select! { - _ = cancel.cancelled() => return, - _ = tokio::time::sleep(period) => {} - } - let Some(inner) = weak.upgrade() else { - return; - }; - let now = Instant::now(); - let mut removed = 0; - inner.entries.retain(|_, entry| { - let retain = !matches!( - entry, - AffinityEntry::Bound { - active_leases: 0, - idle_deadline, - .. - } if *idle_deadline <= now - ); - removed += usize::from(!retain); - retain - }); - inner.entry_count.fetch_sub(removed, Ordering::Relaxed); - } - }); + }) } #[cfg(test)] @@ -177,84 +103,52 @@ impl AffinityCoordinator { request_context: Option<&dyn AsyncEngineContext>, ) -> Result { self.validate_session_id(session_id)?; - let session_id = session_id.as_str().to_string(); - - loop { - let now = Instant::now(); - match self.inner.entries.entry(session_id.clone()) { - Entry::Vacant(entry) => { - self.reserve_entry()?; - return Ok(AffinityAcquire::Initialize(entry.insert_initializing( - &self.inner, - session_id, - requested_target, - ))); + let acquired = if let Some(context) = request_context { + let cancellation = async { + tokio::select! { + biased; + _ = context.stopped() => {}, + _ = context.killed() => {}, + } + }; + let acquired = self + .placement + .acquire_with_cancellation(session_id.as_str(), cancellation) + .await; + match acquired { + Ok(acquired) => acquired, + Err(SessionPlacementError::AcquireCancelled) => { + return Err(cancelled(context.id())); } - Entry::Occupied(mut entry) => match entry.get_mut() { - AffinityEntry::Initializing { notify, .. } => { - #[cfg(test)] - self.inner.waiter_observed.notify_one(); - let notified = notify.clone().notified_owned(); - tokio::pin!(notified); - notified.as_mut().enable(); - drop(entry); - if let Some(context) = request_context { - tokio::select! { - biased; - _ = context.stopped() => { - return Err(cancelled(context.id())); - } - _ = context.killed() => { - return Err(cancelled(context.id())); - } - _ = notified => {} - } - } else { - notified.await; - } - } - AffinityEntry::Bound { - target: _, - revision, - active_leases, - idle_deadline, - } if *active_leases == 0 && *idle_deadline <= now => { - let revision = self.inner.next_revision.fetch_add(1, Ordering::Relaxed); - let notify = Arc::new(Notify::new()); - *entry.get_mut() = AffinityEntry::Initializing { - revision, - notify: notify.clone(), - }; - drop(entry); - return Ok(AffinityAcquire::Initialize(AffinityInitialization { - coordinator: Arc::downgrade(&self.inner), - session_id, - revision, - notify, - requested_target, - active: true, - })); - } - AffinityEntry::Bound { - target, - revision, - active_leases, - .. - } => { - validate_bound_target(&session_id, *target, requested_target)?; - *active_leases += 1; - let lease = AffinityLease { - coordinator: Arc::downgrade(&self.inner), - session_id, - revision: *revision, - active: true, - }; - return Ok(AffinityAcquire::Bound { - target: *target, - lease, - }); - } - }, + Err(error) => return Err(map_placement_error(error)), + } + } else { + self.placement + .acquire(session_id.as_str()) + .await + .map_err(map_placement_error)? + }; + + match acquired { + PlacementAcquire::Initialize(initialization) => { + Ok(AffinityAcquire::Initialize(AffinityInitialization { + initialization, + session_id: session_id.as_str().to_string(), + requested_target, + })) + } + PlacementAcquire::Bound { target, mut lease } => { + let target = *target.target(); + if let Err(error) = + validate_bound_target(session_id.as_str(), target, requested_target) + { + lease.abandon(); + return Err(error); + } + Ok(AffinityAcquire::Bound { + target, + lease: AffinityLease { lease }, + }) } } } @@ -265,119 +159,61 @@ impl AffinityCoordinator { requested_target: Option, ) -> Result, Error> { self.validate_session_id(session_id)?; - let Some(entry) = self.inner.entries.get(session_id.as_str()) else { - return Ok(None); - }; - let AffinityEntry::Bound { - target, - active_leases, - idle_deadline, - .. - } = entry.value() - else { - return Ok(None); - }; - if *active_leases == 0 && *idle_deadline <= Instant::now() { - return Ok(None); + let target = self + .placement + .query(session_id.as_str()) + .map_err(map_placement_error)? + .map(|target| *target.target()); + if let Some(target) = target { + validate_bound_target(session_id.as_str(), target, requested_target)?; + } + Ok(target) + } + + fn validate_session_id(&self, session_id: &SessionAffinityId) -> Result<(), Error> { + if session_id.as_str().len() > self.max_session_id_bytes { + return Err(invalid_argument(format!( + "session affinity ID must not exceed {} bytes", + self.max_session_id_bytes + ))); } - validate_bound_target(session_id.as_str(), *target, requested_target)?; - Ok(Some(*target)) + Ok(()) } #[cfg(test)] pub(super) fn entry_count(&self) -> usize { - self.inner.entry_count.load(Ordering::Relaxed) + self.placement.entry_count() } #[cfg(test)] - pub(super) fn cancellation_token(&self) -> CancellationToken { - self.inner.cancel.clone() + pub(super) fn cancellation_token(&self) -> tokio_util::sync::CancellationToken { + self.placement.cancellation_token() } #[cfg(test)] pub(super) async fn wait_for_reaper(&self) { - self.inner.reaper_started.notified().await; + self.placement.wait_for_reaper().await; + } + + #[cfg(test)] + pub(super) async fn wait_for_reap(&self) { + self.placement.wait_for_reap().await; } #[cfg(test)] pub(super) async fn wait_for_initializing_waiter(&self) { - self.inner.waiter_observed.notified().await; + self.placement.wait_for_initializing_waiter().await; } #[cfg(test)] pub(super) fn expire_for_test(&self, session_id: &SessionAffinityId) { - let Some(mut entry) = self.inner.entries.get_mut(session_id.as_str()) else { - panic!("session affinity entry missing"); - }; - let AffinityEntry::Bound { - active_leases, - idle_deadline, - .. - } = entry.value_mut() - else { - panic!("session affinity entry is not bound"); - }; - assert_eq!(*active_leases, 0); - *idle_deadline = Instant::now(); + self.placement.expire_for_test(session_id.as_str()); } #[cfg(test)] pub(super) fn with_test_limits(max_entries: usize, max_session_id_bytes: usize) -> Self { Self::new_with_limits(Duration::from_secs(10), max_entries, max_session_id_bytes).unwrap() } - - fn validate_session_id(&self, session_id: &SessionAffinityId) -> Result<(), Error> { - if session_id.as_str().len() > self.inner.max_session_id_bytes { - return Err(invalid_argument(format!( - "session affinity ID must not exceed {} bytes", - self.inner.max_session_id_bytes - ))); - } - Ok(()) - } - - fn reserve_entry(&self) -> Result<(), Error> { - self.inner - .entry_count - .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { - (count < self.inner.max_entries).then_some(count + 1) - }) - .map(|_| ()) - .map_err(|_| resource_exhausted("session affinity entry limit reached")) - } -} - -trait VacantEntryExt { - fn insert_initializing( - self, - inner: &Arc, - session_id: String, - requested_target: Option, - ) -> AffinityInitialization; -} - -impl<'a> VacantEntryExt for dashmap::mapref::entry::VacantEntry<'a, String, AffinityEntry> { - fn insert_initializing( - self, - inner: &Arc, - session_id: String, - requested_target: Option, - ) -> AffinityInitialization { - let revision = inner.next_revision.fetch_add(1, Ordering::Relaxed); - let notify = Arc::new(Notify::new()); - self.insert(AffinityEntry::Initializing { - revision, - notify: notify.clone(), - }); - AffinityInitialization { - coordinator: Arc::downgrade(inner), - session_id, - revision, - notify, - requested_target, - active: true, - } - } } pub(crate) enum AffinityAcquire { @@ -424,75 +260,24 @@ impl AffinityAcquire { } pub(crate) struct AffinityInitialization { - coordinator: Weak, + initialization: PlacementInitialization, session_id: String, - revision: u64, - notify: Arc, requested_target: Option, - active: bool, } impl AffinityInitialization { - pub(crate) fn commit(mut self, target: AffinityTarget) -> Result { + pub(crate) fn commit(self, target: AffinityTarget) -> Result { validate_bound_target(&self.session_id, target, self.requested_target)?; - let Some(inner) = self.coordinator.upgrade() else { - return Err(anyhow::anyhow!("session affinity coordinator dropped")); - }; - let Some(mut entry) = inner.entries.get_mut(&self.session_id) else { - return Err(invalid_argument( - "session affinity initialization was cancelled", - )); - }; - if !matches!( - entry.value(), - AffinityEntry::Initializing { revision, .. } if *revision == self.revision - ) { - return Err(invalid_argument("session affinity initialization changed")); - } - *entry = AffinityEntry::Bound { - target, - revision: self.revision, - active_leases: 1, - idle_deadline: Instant::now() + inner.ttl, - }; - drop(entry); - self.active = false; - self.notify.notify_waiters(); - Ok(AffinityLease { - coordinator: Arc::downgrade(&inner), - session_id: self.session_id.clone(), - revision: self.revision, - active: true, - }) - } -} - -impl Drop for AffinityInitialization { - fn drop(&mut self) { - if !self.active { - return; - } - let Some(inner) = self.coordinator.upgrade() else { - return; - }; - let removed = inner.entries.remove_if(&self.session_id, |_, entry| { - matches!( - entry, - AffinityEntry::Initializing { revision, .. } if *revision == self.revision - ) - }); - if removed.is_some() { - inner.entry_count.fetch_sub(1, Ordering::Relaxed); - } - self.notify.notify_waiters(); + let lease = self + .initialization + .commit_already_accepted(target, TargetGeneration::UNVERSIONED) + .map_err(map_placement_error)?; + Ok(AffinityLease { lease }) } } pub(crate) struct AffinityLease { - coordinator: Weak, - session_id: String, - revision: u64, - active: bool, + lease: PlacementLease, } impl AffinityLease { @@ -507,56 +292,8 @@ impl AffinityLease { ) } - fn release(&mut self) { - if !self.active { - return; - } - self.active = false; - let Some(inner) = self.coordinator.upgrade() else { - return; - }; - let Some(mut entry) = inner.entries.get_mut(&self.session_id) else { - return; - }; - let AffinityEntry::Bound { - revision, - active_leases, - idle_deadline, - .. - } = entry.value_mut() - else { - return; - }; - if *revision != self.revision || *active_leases == 0 { - return; - } - *active_leases -= 1; - *idle_deadline = Instant::now() + inner.ttl; - } - fn invalidate(&mut self) { - if !self.active { - return; - } - self.active = false; - let Some(inner) = self.coordinator.upgrade() else { - return; - }; - let removed = inner.entries.remove_if(&self.session_id, |_, entry| { - matches!( - entry, - AffinityEntry::Bound { revision, .. } if *revision == self.revision - ) - }); - if removed.is_some() { - inner.entry_count.fetch_sub(1, Ordering::Relaxed); - } - } -} - -impl Drop for AffinityLease { - fn drop(&mut self) { - self.release(); + self.lease.invalidate(); } } @@ -642,6 +379,37 @@ fn validate_bound_target( } } +fn map_placement_error(error: SessionPlacementError) -> Error { + match error { + SessionPlacementError::KeyTooLong { max_bytes, .. } => invalid_argument(format!( + "session affinity ID must not exceed {max_bytes} bytes" + )), + SessionPlacementError::Capacity { .. } => { + resource_exhausted("session affinity entry limit reached") + } + SessionPlacementError::CoordinatorDropped => { + anyhow::anyhow!("session affinity coordinator dropped") + } + SessionPlacementError::RuntimeUnavailable => { + invalid_argument("session affinity requires a Tokio runtime") + } + SessionPlacementError::InitializationCancelled => { + invalid_argument("session affinity initialization was cancelled") + } + SessionPlacementError::InitializationChanged => { + invalid_argument("session affinity initialization changed") + } + error @ SessionPlacementError::DispatchAmbiguous { .. } + | error @ SessionPlacementError::TargetGenerationChanged { .. } => { + anyhow::Error::new(error) + } + SessionPlacementError::AcquireCancelled => { + invalid_argument("session affinity acquisition was cancelled") + } + SessionPlacementError::InvalidConfig(message) => invalid_argument(message), + } +} + pub(crate) fn invalid_argument(message: impl Into) -> Error { DynamoError::builder() .error_type(ErrorType::InvalidArgument) diff --git a/lib/llm/src/session_affinity/tests.rs b/lib/llm/src/session_affinity/tests.rs index 044eea963356..ff01854b3c92 100644 --- a/lib/llm/src/session_affinity/tests.rs +++ b/lib/llm/src/session_affinity/tests.rs @@ -160,6 +160,26 @@ async fn session_affinity_initialization_is_atomic() { drop(second_lease); } +#[tokio::test(start_paused = true)] +async fn session_affinity_initialization_can_outlive_idle_ttl() { + let coordinator = AffinityCoordinator::new(Duration::from_secs(1)).unwrap(); + let first = coordinator.acquire(&session_id(), None).await.unwrap(); + let AffinityAcquire::Initialize(first) = first else { + panic!("first request must initialize"); + }; + + coordinator.wait_for_reaper().await; + tokio::time::advance(Duration::from_secs(2)).await; + tokio::task::yield_now().await; + + let lease = first.commit(target(7, Some(0))).unwrap(); + assert_eq!( + coordinator.query_target(&session_id(), None).unwrap(), + Some(target(7, Some(0))) + ); + drop(lease); +} + #[tokio::test(start_paused = true)] async fn session_affinity_initializer_cancellation_wakes_waiter() { let coordinator = coordinator(); @@ -213,6 +233,37 @@ async fn session_affinity_wait_stops_when_request_is_cancelled() { drop(first); } +#[tokio::test(start_paused = true)] +async fn session_affinity_only_observes_cancellation_while_waiting() { + let coordinator = coordinator(); + let context = Controller::default(); + context.stop(); + + let first = coordinator + .acquire_with_context(&session_id(), None, &context) + .await + .unwrap(); + let AffinityAcquire::Initialize(first) = first else { + panic!("vacant session must initialize despite a stopped context"); + }; + let first_lease = first.commit(target(7, Some(0))).unwrap(); + + let second = coordinator + .acquire_with_context(&session_id(), None, &context) + .await + .unwrap(); + let AffinityAcquire::Bound { + target: bound_target, + lease: second_lease, + } = second + else { + panic!("bound session must be acquired despite a stopped context"); + }; + assert_eq!(bound_target, target(7, Some(0))); + drop(first_lease); + drop(second_lease); +} + #[tokio::test(start_paused = true)] async fn session_affinity_validates_worker_and_rank_contract() { let coordinator = coordinator(); @@ -438,7 +489,7 @@ async fn session_affinity_reaper_removes_idle_entries_and_stops_on_drop() { coordinator.wait_for_reaper().await; tokio::time::advance(Duration::from_secs(10)).await; - tokio::task::yield_now().await; + coordinator.wait_for_reap().await; assert_eq!(coordinator.entry_count(), 0); drop(coordinator); diff --git a/lib/llm/src/session_placement.rs b/lib/llm/src/session_placement.rs new file mode 100644 index 000000000000..35bc09f25532 --- /dev/null +++ b/lib/llm/src/session_placement.rs @@ -0,0 +1,14 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Process-local coordination for session placement state. + +mod coordinator; + +pub(crate) use coordinator::{ + PlacementAcquire, PlacementInitialization, PlacementLease, SessionPlacement, + SessionPlacementConfig, SessionPlacementError, TargetGeneration, +}; + +#[cfg(test)] +mod tests; diff --git a/lib/llm/src/session_placement/coordinator.rs b/lib/llm/src/session_placement/coordinator.rs new file mode 100644 index 000000000000..e62a6e7027c3 --- /dev/null +++ b/lib/llm/src/session_placement/coordinator.rs @@ -0,0 +1,1261 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::{ + future::Future, + sync::{ + Arc, Weak, + atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use dashmap::{DashMap, mapref::entry::Entry}; +use parking_lot::Mutex; +use thiserror::Error; +use tokio::{runtime::Handle, sync::Notify, task::JoinHandle, time::Instant}; +use tokio_util::sync::CancellationToken; + +const MIN_IDLE_TTL: Duration = Duration::from_secs(1); +const MAX_IDLE_TTL: Duration = Duration::from_secs(31_536_000); +// DashMap removal and the global count update are separate atomic operations. Retry a bounded +// number of times when a concurrent removal has exposed a map slot but not yet released its count. +const CAPACITY_CONTENTION_RETRIES: usize = 2; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct SessionPlacementConfig { + pub(crate) idle_ttl: Duration, + pub(crate) initialization_timeout: Option, + pub(crate) max_entries: usize, + pub(crate) max_key_bytes: usize, +} + +#[derive(Debug, Error, PartialEq, Eq)] +pub(crate) enum SessionPlacementError { + #[error("invalid session placement configuration: {0}")] + InvalidConfig(&'static str), + + #[error("session placement key is {actual_bytes} bytes; maximum is {max_bytes}")] + KeyTooLong { + actual_bytes: usize, + max_bytes: usize, + }, + + #[error("session placement entry limit of {max_entries} reached")] + Capacity { max_entries: usize }, + + #[error("session placement coordinator was dropped")] + CoordinatorDropped, + + #[error("session placement requires a Tokio runtime")] + RuntimeUnavailable, + + #[error("session placement initialization was cancelled")] + InitializationCancelled, + + #[error("session placement initialization changed")] + InitializationChanged, + + #[error( + "session placement dispatch {attempt_id} for target generation {target_generation} has an ambiguous outcome" + )] + DispatchAmbiguous { + attempt_id: u64, + target_generation: u64, + }, + + #[error( + "session placement target generation changed from {expected_generation} to {actual_generation}" + )] + TargetGenerationChanged { + expected_generation: u64, + actual_generation: u64, + }, + + #[error("session placement acquisition was cancelled")] + AcquireCancelled, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub(crate) struct PlacementAttemptId(u64); + +impl PlacementAttemptId { + pub(crate) fn get(self) -> u64 { + self.0 + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub(crate) struct TargetGeneration(u64); + +impl TargetGeneration { + pub(crate) const UNVERSIONED: Self = Self(0); + + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "versioned targets are consumed by the global-router integration" + ) + )] + pub(crate) fn new(generation: u64) -> Self { + Self(generation) + } + + pub(crate) fn get(self) -> u64 { + self.0 + } +} + +struct VersionedTargetInner { + target: T, + generation: TargetGeneration, +} + +pub(crate) struct VersionedTarget { + inner: Arc>, +} + +impl VersionedTarget { + fn new(target: T, generation: TargetGeneration) -> Self { + Self { + inner: Arc::new(VersionedTargetInner { target, generation }), + } + } + + pub(crate) fn target(&self) -> &T { + &self.inner.target + } + + pub(crate) fn generation(&self) -> TargetGeneration { + self.inner.generation + } + + fn same_as(&self, other: &Self) -> bool { + self.generation() == other.generation() && Arc::ptr_eq(&self.inner, &other.inner) + } +} + +impl Clone for VersionedTarget { + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + } + } +} + +// A reservation may be replaced before dispatch begins. Once a target is Dispatching, every +// non-definitive outcome is quarantined as Ambiguous so another owner cannot silently replay it. +enum PlacementEntry { + Reserved { + attempt: PlacementAttemptId, + notify: Arc, + deadline: Option, + }, + Dispatching { + attempt: PlacementAttemptId, + candidate: VersionedTarget, + notify: Arc, + deadline: Option, + }, + Ambiguous { + attempt: PlacementAttemptId, + candidate: VersionedTarget, + }, + Bound { + target: VersionedTarget, + revision: PlacementAttemptId, + active_leases: usize, + idle_deadline: Instant, + }, +} + +struct SessionPlacementInner { + entries: DashMap>, + config: SessionPlacementConfig, + entry_count: AtomicUsize, + next_attempt: AtomicU64, + cancel: CancellationToken, + reaper_running: AtomicBool, + reaper: Mutex>>, + #[cfg(test)] + reaper_started: Arc, + #[cfg(test)] + reaper_completed: Arc, + #[cfg(test)] + waiter_observed: Arc, +} + +struct ReaperRunning { + inner: Weak>, +} + +impl Drop for ReaperRunning { + fn drop(&mut self) { + if let Some(inner) = self.inner.upgrade() { + inner.reaper_running.store(false, Ordering::Release); + } + } +} + +impl Drop for SessionPlacementInner { + fn drop(&mut self) { + self.cancel.cancel(); + if let Some(reaper) = self.reaper.get_mut().take() { + reaper.abort(); + } + } +} + +pub(crate) struct SessionPlacement { + inner: Arc>, +} + +impl Clone for SessionPlacement { + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + } + } +} + +impl SessionPlacement +where + T: Send + Sync + 'static, +{ + pub(crate) fn new(config: SessionPlacementConfig) -> Result { + validate_config(config)?; + let handle = + Handle::try_current().map_err(|_| SessionPlacementError::RuntimeUnavailable)?; + let inner = Arc::new(SessionPlacementInner { + entries: DashMap::new(), + config, + entry_count: AtomicUsize::new(0), + next_attempt: AtomicU64::new(1), + cancel: CancellationToken::new(), + reaper_running: AtomicBool::new(false), + reaper: Mutex::new(None), + #[cfg(test)] + reaper_started: Arc::new(Notify::new()), + #[cfg(test)] + reaper_completed: Arc::new(Notify::new()), + #[cfg(test)] + waiter_observed: Arc::new(Notify::new()), + }); + *inner.reaper.lock() = Some(Self::spawn_reaper_task(&inner, &handle)); + Ok(Self { inner }) + } + + pub(crate) async fn acquire( + &self, + key: &str, + ) -> Result, SessionPlacementError> { + self.acquire_inner(key, std::future::pending()).await + } + + pub(crate) async fn acquire_with_cancellation( + &self, + key: &str, + cancellation: F, + ) -> Result, SessionPlacementError> + where + F: Future, + { + self.acquire_inner(key, cancellation).await + } + + async fn acquire_inner( + &self, + key: &str, + cancellation: F, + ) -> Result, SessionPlacementError> + where + F: Future, + { + self.validate_key(key)?; + self.ensure_reaper(); + let key = key.to_owned(); + let mut reaped_for_capacity = false; + let mut capacity_contention_retries = 0; + tokio::pin!(cancellation); + + loop { + let now = Instant::now(); + match self.inner.entries.entry(key.clone()) { + Entry::Vacant(entry) => { + if let Err(error) = self.reserve_entry() { + drop(entry); + if reaped_for_capacity { + if capacity_contention_retries < CAPACITY_CONTENTION_RETRIES + && self.inner.entries.len() < self.inner.config.max_entries + { + capacity_contention_retries += 1; + tokio::task::yield_now().await; + continue; + } + return Err(error); + } + self.reap_expired_async().await; + reaped_for_capacity = true; + continue; + } + let attempt = self.next_attempt(); + let notify = Arc::new(Notify::new()); + entry.insert(PlacementEntry::Reserved { + attempt, + notify: notify.clone(), + deadline: self.initialization_deadline(now), + }); + return Ok(PlacementAcquire::Initialize(PlacementInitialization { + placement: Arc::downgrade(&self.inner), + key, + attempt, + notify, + active: true, + })); + } + Entry::Occupied(mut entry) => match entry.get_mut() { + PlacementEntry::Reserved { + notify, deadline, .. + } if deadline.as_ref().is_some_and(|deadline| *deadline <= now) => { + let stale_notify = notify.clone(); + let attempt = self.next_attempt(); + let notify = Arc::new(Notify::new()); + *entry.get_mut() = PlacementEntry::Reserved { + attempt, + notify: notify.clone(), + deadline: self.initialization_deadline(now), + }; + drop(entry); + stale_notify.notify_waiters(); + return Ok(PlacementAcquire::Initialize(PlacementInitialization { + placement: Arc::downgrade(&self.inner), + key, + attempt, + notify, + active: true, + })); + } + PlacementEntry::Reserved { + notify, deadline, .. + } => { + let deadline = *deadline; + let notified = notify.clone().notified_owned(); + self.wait_for_initialization( + entry, + notified, + deadline, + cancellation.as_mut(), + ) + .await?; + } + PlacementEntry::Dispatching { + attempt, + candidate, + notify, + deadline, + } if deadline.as_ref().is_some_and(|deadline| *deadline <= now) => { + let attempt = *attempt; + let candidate = candidate.clone(); + let notify = notify.clone(); + *entry.get_mut() = PlacementEntry::Ambiguous { + attempt, + candidate: candidate.clone(), + }; + drop(entry); + notify.notify_waiters(); + return Err(dispatch_ambiguous(attempt, candidate.generation())); + } + PlacementEntry::Dispatching { + notify, deadline, .. + } => { + let deadline = *deadline; + let notified = notify.clone().notified_owned(); + self.wait_for_initialization( + entry, + notified, + deadline, + cancellation.as_mut(), + ) + .await?; + } + PlacementEntry::Ambiguous { attempt, candidate } => { + return Err(dispatch_ambiguous(*attempt, candidate.generation())); + } + PlacementEntry::Bound { + target, + active_leases, + idle_deadline, + .. + } if *active_leases == 0 && *idle_deadline <= now => { + let stale_target = target.clone(); + let attempt = self.next_attempt(); + let notify = Arc::new(Notify::new()); + *entry.get_mut() = PlacementEntry::Reserved { + attempt, + notify: notify.clone(), + deadline: self.initialization_deadline(now), + }; + drop(entry); + drop(stale_target); + return Ok(PlacementAcquire::Initialize(PlacementInitialization { + placement: Arc::downgrade(&self.inner), + key, + attempt, + notify, + active: true, + })); + } + PlacementEntry::Bound { + target, + revision, + active_leases, + .. + } => { + *active_leases += 1; + return Ok(PlacementAcquire::Bound { + target: target.clone(), + lease: PlacementLease { + placement: Arc::downgrade(&self.inner), + key, + revision: *revision, + active: true, + }, + }); + } + }, + } + } + } + + async fn wait_for_initialization( + &self, + entry: dashmap::mapref::entry::OccupiedEntry<'_, String, PlacementEntry>, + notified: tokio::sync::futures::OwnedNotified, + deadline: Option, + mut cancellation: std::pin::Pin<&mut F>, + ) -> Result<(), SessionPlacementError> + where + F: Future, + { + #[cfg(test)] + self.inner.waiter_observed.notify_one(); + tokio::pin!(notified); + notified.as_mut().enable(); + drop(entry); + let timeout = async move { + match deadline { + Some(deadline) => tokio::time::sleep_until(deadline).await, + None => std::future::pending::<()>().await, + } + }; + tokio::pin!(timeout); + tokio::select! { + biased; + _ = cancellation.as_mut() => Err(SessionPlacementError::AcquireCancelled), + _ = notified => Ok(()), + _ = timeout.as_mut() => Ok(()), + } + } + + pub(crate) fn query( + &self, + key: &str, + ) -> Result>, SessionPlacementError> { + self.validate_key(key)?; + self.ensure_reaper(); + let Some(entry) = self.inner.entries.get(key) else { + return Ok(None); + }; + match entry.value() { + PlacementEntry::Ambiguous { attempt, candidate } => { + Err(dispatch_ambiguous(*attempt, candidate.generation())) + } + PlacementEntry::Bound { + target, + active_leases, + idle_deadline, + .. + } if *active_leases > 0 || *idle_deadline > Instant::now() => Ok(Some(target.clone())), + _ => Ok(None), + } + } + + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "ambiguous dispatch recovery is consumed by the global-router integration" + ) + )] + pub(crate) fn resolve_ambiguous( + &self, + key: &str, + attempt: PlacementAttemptId, + generation: TargetGeneration, + resolution: AmbiguousResolution, + ) -> Result>, SessionPlacementError> { + self.validate_key(key)?; + match resolution { + AmbiguousResolution::Accepted => { + let Some(mut entry) = self.inner.entries.get_mut(key) else { + return Err(SessionPlacementError::InitializationCancelled); + }; + let PlacementEntry::Ambiguous { + attempt: current, + candidate, + } = entry.value() + else { + return Err(SessionPlacementError::InitializationChanged); + }; + validate_attempt(*current, candidate, attempt, generation)?; + let target = candidate.clone(); + *entry = PlacementEntry::Bound { + target, + revision: attempt, + active_leases: 1, + idle_deadline: Instant::now() + self.inner.config.idle_ttl, + }; + drop(entry); + Ok(Some(PlacementLease { + placement: Arc::downgrade(&self.inner), + key: key.to_owned(), + revision: attempt, + active: true, + })) + } + AmbiguousResolution::DefinitelyNotAccepted => { + let removed = self.inner.entries.remove_if(key, |_, entry| { + matches!( + entry, + PlacementEntry::Ambiguous { + attempt: current, + candidate, + } if *current == attempt && candidate.generation() == generation + ) + }); + if removed.is_none() { + return self.ambiguous_resolution_error(key, attempt, generation); + } + self.inner.entry_count.fetch_sub(1, Ordering::Relaxed); + drop(removed); + Ok(None) + } + } + } + + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "only used by the deferred global-router recovery path" + ) + )] + fn ambiguous_resolution_error( + &self, + key: &str, + attempt: PlacementAttemptId, + generation: TargetGeneration, + ) -> Result>, SessionPlacementError> { + let Some(entry) = self.inner.entries.get(key) else { + return Err(SessionPlacementError::InitializationCancelled); + }; + let PlacementEntry::Ambiguous { + attempt: current, + candidate, + } = entry.value() + else { + return Err(SessionPlacementError::InitializationChanged); + }; + validate_attempt(*current, candidate, attempt, generation)?; + Err(SessionPlacementError::InitializationChanged) + } + + fn spawn_reaper_task(inner: &Arc>, handle: &Handle) -> JoinHandle<()> { + let weak = Arc::downgrade(inner); + let running = ReaperRunning { + inner: weak.clone(), + }; + inner.reaper_running.store(true, Ordering::Release); + let cancel = inner.cancel.clone(); + let period = inner + .config + .initialization_timeout + .map_or(inner.config.idle_ttl, |timeout| { + inner.config.idle_ttl.min(timeout) + }) + .min(Duration::from_secs(30)); + #[cfg(test)] + let reaper_started = inner.reaper_started.clone(); + #[cfg(test)] + let reaper_completed = inner.reaper_completed.clone(); + handle.spawn(async move { + let _running = running; + let sleep = tokio::time::sleep(period); + tokio::pin!(sleep); + #[cfg(test)] + reaper_started.notify_one(); + loop { + tokio::select! { + _ = cancel.cancelled() => return, + _ = sleep.as_mut() => {} + } + let Some(inner) = weak.upgrade() else { + return; + }; + let cleanup = tokio::task::spawn_blocking(move || { + inner.reap_expired(Instant::now()); + }); + match cleanup.await { + Ok(()) => { + #[cfg(test)] + reaper_completed.notify_one(); + } + Err(error) => { + tracing::warn!(?error, "session placement cleanup task failed"); + } + } + sleep.as_mut().reset(Instant::now() + period); + } + }) + } + + fn ensure_reaper(&self) { + if self.inner.reaper_running.load(Ordering::Acquire) { + return; + } + let mut reaper = self.inner.reaper.lock(); + if self.inner.reaper_running.load(Ordering::Acquire) { + return; + } + if let Ok(handle) = Handle::try_current() { + *reaper = Some(Self::spawn_reaper_task(&self.inner, &handle)); + } + } + + async fn reap_expired_async(&self) { + let inner = self.inner.clone(); + if let Err(error) = + tokio::task::spawn_blocking(move || inner.reap_expired(Instant::now())).await + { + tracing::warn!(?error, "session placement capacity cleanup task failed"); + } + } + + fn validate_key(&self, key: &str) -> Result<(), SessionPlacementError> { + if key.len() > self.inner.config.max_key_bytes { + return Err(SessionPlacementError::KeyTooLong { + actual_bytes: key.len(), + max_bytes: self.inner.config.max_key_bytes, + }); + } + Ok(()) + } + + fn reserve_entry(&self) -> Result<(), SessionPlacementError> { + self.inner + .entry_count + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |count| { + (count < self.inner.config.max_entries).then_some(count + 1) + }) + .map(|_| ()) + .map_err(|_| SessionPlacementError::Capacity { + max_entries: self.inner.config.max_entries, + }) + } + + fn next_attempt(&self) -> PlacementAttemptId { + PlacementAttemptId(self.inner.next_attempt.fetch_add(1, Ordering::Relaxed)) + } + + fn initialization_deadline(&self, now: Instant) -> Option { + self.inner + .config + .initialization_timeout + .map(|timeout| now + timeout) + } + + #[cfg(test)] + pub(crate) fn entry_count(&self) -> usize { + self.inner.entry_count.load(Ordering::Relaxed) + } + + #[cfg(test)] + pub(crate) fn cancellation_token(&self) -> CancellationToken { + self.inner.cancel.clone() + } + + #[cfg(test)] + pub(crate) async fn wait_for_reaper(&self) { + self.inner.reaper_started.notified().await; + } + + #[cfg(test)] + pub(crate) async fn wait_for_reap(&self) { + self.inner.reaper_completed.notified().await; + } + + #[cfg(test)] + pub(crate) async fn stop_reaper_for_test(&self) { + let reaper = { + let mut reaper = self.inner.reaper.lock(); + reaper.take() + }; + if let Some(reaper) = reaper { + reaper.abort(); + let _ = reaper.await; + } + } + + #[cfg(test)] + pub(crate) async fn wait_for_initializing_waiter(&self) { + self.inner.waiter_observed.notified().await; + } + + #[cfg(test)] + pub(crate) fn expire_for_test(&self, key: &str) { + let Some(mut entry) = self.inner.entries.get_mut(key) else { + panic!("session placement entry missing"); + }; + let PlacementEntry::Bound { + active_leases, + idle_deadline, + .. + } = entry.value_mut() + else { + panic!("session placement entry is not bound"); + }; + assert_eq!(*active_leases, 0); + *idle_deadline = Instant::now(); + } +} + +impl SessionPlacementInner +where + T: Send + Sync + 'static, +{ + fn reap_expired(&self, now: Instant) { + let keys: Vec = self + .entries + .iter() + .filter_map(|entry| { + let is_expired = match entry.value() { + PlacementEntry::Reserved { deadline, .. } + | PlacementEntry::Dispatching { deadline, .. } => { + deadline.as_ref().is_some_and(|deadline| *deadline <= now) + } + PlacementEntry::Bound { + active_leases: 0, + idle_deadline, + .. + } => *idle_deadline <= now, + _ => false, + }; + is_expired.then(|| entry.key().clone()) + }) + .collect(); + + for key in keys { + let mut transitioned_notify = None; + if let Some(mut entry) = self.entries.get_mut(&key) + && let PlacementEntry::Dispatching { + attempt, + candidate, + notify, + deadline, + } = entry.value_mut() + && deadline.as_ref().is_some_and(|deadline| *deadline <= now) + { + let attempt = *attempt; + let candidate = candidate.clone(); + transitioned_notify = Some(notify.clone()); + *entry = PlacementEntry::Ambiguous { attempt, candidate }; + } + if let Some(notify) = transitioned_notify { + notify.notify_waiters(); + continue; + } + + let removed = self.entries.remove_if(&key, |_, entry| { + matches!( + entry, + PlacementEntry::Reserved { + deadline: Some(deadline), + .. + } if *deadline <= now + ) || matches!( + entry, + PlacementEntry::Bound { + active_leases: 0, + idle_deadline, + .. + } if *idle_deadline <= now + ) + }); + let Some((_key, removed_entry)) = removed else { + continue; + }; + self.entry_count.fetch_sub(1, Ordering::Relaxed); + if let PlacementEntry::Reserved { notify, .. } = &removed_entry { + notify.notify_waiters(); + } + drop(removed_entry); + } + } +} + +pub(crate) enum PlacementAcquire { + Initialize(PlacementInitialization), + Bound { + target: VersionedTarget, + lease: PlacementLease, + }, +} + +pub(crate) struct PlacementInitialization { + placement: Weak>, + key: String, + attempt: PlacementAttemptId, + notify: Arc, + active: bool, +} + +impl PlacementInitialization +where + T: Send + Sync + 'static, +{ + pub(crate) fn begin_dispatch( + mut self, + target: T, + generation: TargetGeneration, + ) -> Result, SessionPlacementError> { + let Some(inner) = self.placement.upgrade() else { + return Err(SessionPlacementError::CoordinatorDropped); + }; + let Some(mut entry) = inner.entries.get_mut(&self.key) else { + return Err(SessionPlacementError::InitializationCancelled); + }; + let PlacementEntry::Reserved { + attempt, + notify, + deadline, + } = entry.value() + else { + return Err(SessionPlacementError::InitializationChanged); + }; + if *attempt != self.attempt { + return Err(SessionPlacementError::InitializationChanged); + } + if deadline + .as_ref() + .is_some_and(|deadline| *deadline <= Instant::now()) + { + drop(entry); + return Err(SessionPlacementError::InitializationCancelled); + } + let notify = notify.clone(); + let deadline = *deadline; + let candidate = VersionedTarget::new(target, generation); + *entry = PlacementEntry::Dispatching { + attempt: self.attempt, + candidate: candidate.clone(), + notify, + deadline, + }; + drop(entry); + self.active = false; + Ok(PlacementDispatch { + placement: Arc::downgrade(&inner), + key: self.key.clone(), + attempt: self.attempt, + candidate, + notify: self.notify.clone(), + active: true, + }) + } + + pub(crate) fn commit_already_accepted( + self, + target: T, + generation: TargetGeneration, + ) -> Result, SessionPlacementError> { + // Compatibility only: callers must not use this around an in-flight dispatch. New + // forwarding paths must call `begin_dispatch` before sending and finish with the + // transport's explicit outcome. + self.begin_dispatch(target, generation)? + .finish(PlacementDispatchOutcome::Accepted)? + .ok_or(SessionPlacementError::InitializationChanged) + } +} + +impl Drop for PlacementInitialization { + fn drop(&mut self) { + if !self.active { + return; + } + let Some(inner) = self.placement.upgrade() else { + return; + }; + let removed = inner.entries.remove_if(&self.key, |_, entry| { + matches!( + entry, + PlacementEntry::Reserved { attempt, .. } if *attempt == self.attempt + ) + }); + if removed.is_some() { + inner.entry_count.fetch_sub(1, Ordering::Relaxed); + } + drop(removed); + self.notify.notify_waiters(); + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum PlacementDispatchOutcome { + Accepted, + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "transport outcome mapping is added with the global-router integration" + ) + )] + DefinitelyNotAccepted, + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "transport outcome mapping is added with the global-router integration" + ) + )] + Ambiguous, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[cfg_attr( + not(test), + allow( + dead_code, + reason = "ambiguous dispatch recovery is consumed by the global-router integration" + ) +)] +pub(crate) enum AmbiguousResolution { + Accepted, + DefinitelyNotAccepted, +} + +pub(crate) struct PlacementDispatch { + placement: Weak>, + key: String, + attempt: PlacementAttemptId, + candidate: VersionedTarget, + notify: Arc, + active: bool, +} + +impl PlacementDispatch +where + T: Send + Sync + 'static, +{ + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "attempt IDs are consumed by the global-router recovery path" + ) + )] + pub(crate) fn attempt_id(&self) -> PlacementAttemptId { + self.attempt + } + + #[cfg_attr( + not(test), + allow( + dead_code, + reason = "the global-router dispatch path reads the reserved target" + ) + )] + pub(crate) fn target(&self) -> &VersionedTarget { + &self.candidate + } + + pub(crate) fn finish( + mut self, + outcome: PlacementDispatchOutcome, + ) -> Result>, SessionPlacementError> { + let Some(inner) = self.placement.upgrade() else { + return Err(SessionPlacementError::CoordinatorDropped); + }; + match outcome { + PlacementDispatchOutcome::Accepted => { + let Some(mut entry) = inner.entries.get_mut(&self.key) else { + return Err(SessionPlacementError::InitializationCancelled); + }; + let notify = match entry.value() { + PlacementEntry::Dispatching { + attempt, + candidate, + notify, + .. + } => { + validate_attempt( + *attempt, + candidate, + self.attempt, + self.candidate.generation(), + )?; + Some(notify.clone()) + } + PlacementEntry::Ambiguous { attempt, candidate } => { + validate_attempt( + *attempt, + candidate, + self.attempt, + self.candidate.generation(), + )?; + None + } + _ => return Err(SessionPlacementError::InitializationChanged), + }; + *entry = PlacementEntry::Bound { + target: self.candidate.clone(), + revision: self.attempt, + active_leases: 1, + idle_deadline: Instant::now() + inner.config.idle_ttl, + }; + drop(entry); + self.active = false; + if let Some(notify) = notify { + notify.notify_waiters(); + } + Ok(Some(PlacementLease { + placement: Arc::downgrade(&inner), + key: self.key.clone(), + revision: self.attempt, + active: true, + })) + } + PlacementDispatchOutcome::DefinitelyNotAccepted => { + let removed = inner.entries.remove_if(&self.key, |_, entry| { + matches!( + entry, + PlacementEntry::Dispatching { + attempt, + candidate, + .. + } | PlacementEntry::Ambiguous { + attempt, + candidate, + } if *attempt == self.attempt && candidate.same_as(&self.candidate) + ) + }); + if removed.is_none() { + return Err(SessionPlacementError::InitializationChanged); + } + inner.entry_count.fetch_sub(1, Ordering::Relaxed); + let notify = removed.as_ref().and_then(|(_, entry)| match entry { + PlacementEntry::Dispatching { notify, .. } => Some(notify.clone()), + _ => None, + }); + drop(removed); + self.active = false; + if let Some(notify) = notify { + notify.notify_waiters(); + } + Ok(None) + } + PlacementDispatchOutcome::Ambiguous => { + self.mark_ambiguous(&inner)?; + self.active = false; + Ok(None) + } + } + } + + fn mark_ambiguous( + &self, + inner: &Arc>, + ) -> Result<(), SessionPlacementError> { + let Some(mut entry) = inner.entries.get_mut(&self.key) else { + return Err(SessionPlacementError::InitializationCancelled); + }; + match entry.value() { + PlacementEntry::Dispatching { + attempt, candidate, .. + } => validate_attempt( + *attempt, + candidate, + self.attempt, + self.candidate.generation(), + )?, + PlacementEntry::Ambiguous { attempt, candidate } => { + return validate_attempt( + *attempt, + candidate, + self.attempt, + self.candidate.generation(), + ); + } + _ => return Err(SessionPlacementError::InitializationChanged), + } + *entry = PlacementEntry::Ambiguous { + attempt: self.attempt, + candidate: self.candidate.clone(), + }; + drop(entry); + self.notify.notify_waiters(); + Ok(()) + } +} + +impl Drop for PlacementDispatch { + fn drop(&mut self) { + if !self.active { + return; + } + let Some(inner) = self.placement.upgrade() else { + return; + }; + let Some(mut entry) = inner.entries.get_mut(&self.key) else { + return; + }; + let PlacementEntry::Dispatching { + attempt, candidate, .. + } = entry.value() + else { + return; + }; + if *attempt != self.attempt || !candidate.same_as(&self.candidate) { + return; + } + *entry = PlacementEntry::Ambiguous { + attempt: self.attempt, + candidate: self.candidate.clone(), + }; + drop(entry); + self.notify.notify_waiters(); + } +} + +pub(crate) struct PlacementLease { + placement: Weak>, + key: String, + revision: PlacementAttemptId, + active: bool, +} + +impl PlacementLease { + pub(crate) fn invalidate(&mut self) { + if !self.active { + return; + } + self.active = false; + let Some(inner) = self.placement.upgrade() else { + return; + }; + let removed = inner.entries.remove_if(&self.key, |_, entry| { + matches!( + entry, + PlacementEntry::Bound { revision, .. } if *revision == self.revision + ) + }); + if removed.is_some() { + inner.entry_count.fetch_sub(1, Ordering::Relaxed); + } + drop(removed); + } + + pub(crate) fn abandon(&mut self) { + self.release(false); + } + + fn release(&mut self, refresh_ttl: bool) { + if !self.active { + return; + } + self.active = false; + let Some(inner) = self.placement.upgrade() else { + return; + }; + let Some(mut entry) = inner.entries.get_mut(&self.key) else { + return; + }; + let PlacementEntry::Bound { + revision, + active_leases, + idle_deadline, + .. + } = entry.value_mut() + else { + return; + }; + if *revision != self.revision || *active_leases == 0 { + return; + } + *active_leases -= 1; + if refresh_ttl { + *idle_deadline = Instant::now() + inner.config.idle_ttl; + } + } +} + +impl Drop for PlacementLease { + fn drop(&mut self) { + self.release(true); + } +} + +fn validate_attempt( + current_attempt: PlacementAttemptId, + current_target: &VersionedTarget, + expected_attempt: PlacementAttemptId, + expected_generation: TargetGeneration, +) -> Result<(), SessionPlacementError> { + if current_attempt != expected_attempt { + return Err(SessionPlacementError::InitializationChanged); + } + let actual_generation = current_target.generation(); + if actual_generation != expected_generation { + return Err(SessionPlacementError::TargetGenerationChanged { + expected_generation: expected_generation.get(), + actual_generation: actual_generation.get(), + }); + } + Ok(()) +} + +fn dispatch_ambiguous( + attempt: PlacementAttemptId, + generation: TargetGeneration, +) -> SessionPlacementError { + SessionPlacementError::DispatchAmbiguous { + attempt_id: attempt.get(), + target_generation: generation.get(), + } +} + +fn validate_config(config: SessionPlacementConfig) -> Result<(), SessionPlacementError> { + if !(MIN_IDLE_TTL..=MAX_IDLE_TTL).contains(&config.idle_ttl) { + return Err(SessionPlacementError::InvalidConfig( + "idle_ttl must be between 1 second and 1 year", + )); + } + if let Some(initialization_timeout) = config.initialization_timeout + && !(MIN_IDLE_TTL..=MAX_IDLE_TTL).contains(&initialization_timeout) + { + return Err(SessionPlacementError::InvalidConfig( + "initialization_timeout must be between 1 second and 1 year when configured", + )); + } + if config.max_entries == 0 { + return Err(SessionPlacementError::InvalidConfig( + "max_entries must be greater than zero", + )); + } + if config.max_key_bytes == 0 { + return Err(SessionPlacementError::InvalidConfig( + "max_key_bytes must be greater than zero", + )); + } + Ok(()) +} diff --git a/lib/llm/src/session_placement/tests.rs b/lib/llm/src/session_placement/tests.rs new file mode 100644 index 000000000000..9bc46ab6307b --- /dev/null +++ b/lib/llm/src/session_placement/tests.rs @@ -0,0 +1,583 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::{sync::Arc, time::Duration}; + +use super::coordinator::{ + AmbiguousResolution, PlacementAcquire, PlacementDispatch, PlacementDispatchOutcome, + PlacementInitialization, PlacementLease, SessionPlacement, SessionPlacementConfig, + SessionPlacementError, TargetGeneration, +}; + +fn placement(max_entries: usize, max_key_bytes: usize) -> SessionPlacement +where + T: Send + Sync + 'static, +{ + SessionPlacement::new(SessionPlacementConfig { + idle_ttl: Duration::from_secs(10), + initialization_timeout: Some(Duration::from_secs(10)), + max_entries, + max_key_bytes, + }) + .unwrap() +} + +fn begin( + initialization: PlacementInitialization, + target: T, + generation: u64, +) -> PlacementDispatch +where + T: Send + Sync + 'static, +{ + initialization + .begin_dispatch(target, TargetGeneration::new(generation)) + .unwrap() +} + +fn accept( + initialization: PlacementInitialization, + target: T, + generation: u64, +) -> PlacementLease +where + T: Send + Sync + 'static, +{ + begin(initialization, target, generation) + .finish(PlacementDispatchOutcome::Accepted) + .unwrap() + .unwrap() +} + +fn query_string(placement: &SessionPlacement, key: &str) -> Option { + placement + .query(key) + .unwrap() + .map(|target| target.target().clone()) +} + +#[tokio::test] +async fn concurrent_miss_initializes_once() { + let placement = Arc::new(placement(10, 32)); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + let dispatch = begin(initialization, "cluster-a".to_string(), 7); + + let waiter = { + let placement = placement.clone(); + tokio::spawn(async move { placement.acquire("session").await }) + }; + placement.wait_for_initializing_waiter().await; + + let first_lease = dispatch + .finish(PlacementDispatchOutcome::Accepted) + .unwrap() + .unwrap(); + let second = waiter.await.unwrap().unwrap(); + let PlacementAcquire::Bound { + target, + lease: second_lease, + } = second + else { + panic!("waiter should observe the committed placement"); + }; + assert_eq!(target.target(), "cluster-a"); + assert_eq!(target.generation(), TargetGeneration::new(7)); + assert_eq!(placement.entry_count(), 1); + + drop(first_lease); + drop(second_lease); +} + +#[tokio::test] +async fn dropped_reservation_rolls_back_and_wakes_waiter() { + let placement = Arc::new(placement::(10, 32)); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + + let waiter = { + let placement = placement.clone(); + tokio::spawn(async move { placement.acquire("session").await }) + }; + placement.wait_for_initializing_waiter().await; + drop(initialization); + + let next = waiter.await.unwrap().unwrap(); + let PlacementAcquire::Initialize(initialization) = next else { + panic!("waiter should become the next initializer"); + }; + let lease = accept(initialization, "cluster-b".to_string(), 1); + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-b") + ); + assert_eq!(placement.entry_count(), 1); + drop(lease); +} + +#[tokio::test(start_paused = true)] +async fn dropped_dispatch_is_quarantined_as_ambiguous() { + let placement = Arc::new(placement(10, 32)); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + let dispatch = begin(initialization, "cluster-a".to_string(), 3); + + let waiter = { + let placement = placement.clone(); + tokio::spawn(async move { placement.acquire("session").await }) + }; + placement.wait_for_initializing_waiter().await; + drop(dispatch); + + assert!(matches!( + waiter.await.unwrap(), + Err(SessionPlacementError::DispatchAmbiguous { + target_generation: 3, + .. + }) + )); + assert!(matches!( + placement.query("session"), + Err(SessionPlacementError::DispatchAmbiguous { + target_generation: 3, + .. + }) + )); + + placement.wait_for_reaper().await; + tokio::time::advance(Duration::from_secs(30)).await; + tokio::task::yield_now().await; + assert_eq!(placement.entry_count(), 1); +} + +#[tokio::test] +async fn definitely_not_accepted_allows_one_waiter_to_retry() { + let placement = Arc::new(placement(10, 32)); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + let dispatch = begin(initialization, "cluster-a".to_string(), 1); + + let waiter = { + let placement = placement.clone(); + tokio::spawn(async move { placement.acquire("session").await }) + }; + placement.wait_for_initializing_waiter().await; + assert!( + dispatch + .finish(PlacementDispatchOutcome::DefinitelyNotAccepted) + .unwrap() + .is_none() + ); + + let retry = waiter.await.unwrap().unwrap(); + assert!(matches!(&retry, PlacementAcquire::Initialize(_))); + drop(retry); + assert_eq!(placement.entry_count(), 0); +} + +#[tokio::test(start_paused = true)] +async fn dispatch_timeout_blocks_replay_but_late_acceptance_can_commit() { + let placement = Arc::new(placement(10, 32)); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + let dispatch = begin(initialization, "cluster-a".to_string(), 9); + + let waiter = { + let placement = placement.clone(); + tokio::spawn(async move { placement.acquire("session").await }) + }; + placement.wait_for_initializing_waiter().await; + tokio::time::advance(Duration::from_secs(11)).await; + tokio::task::yield_now().await; + + assert!(matches!( + waiter.await.unwrap(), + Err(SessionPlacementError::DispatchAmbiguous { + target_generation: 9, + .. + }) + )); + let lease = dispatch + .finish(PlacementDispatchOutcome::Accepted) + .unwrap() + .unwrap(); + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-a") + ); + drop(lease); +} + +#[tokio::test] +async fn ambiguous_resolution_is_fenced_by_attempt_and_generation() { + let placement = placement(1, 32); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + let dispatch = begin(initialization, "cluster-a".to_string(), 4); + let attempt = dispatch.attempt_id(); + assert_eq!(dispatch.target().target(), "cluster-a"); + drop(dispatch); + + assert_eq!( + placement + .resolve_ambiguous( + "session", + attempt, + TargetGeneration::new(5), + AmbiguousResolution::DefinitelyNotAccepted, + ) + .err() + .unwrap(), + SessionPlacementError::TargetGenerationChanged { + expected_generation: 5, + actual_generation: 4, + } + ); + assert!( + placement + .resolve_ambiguous( + "session", + attempt, + TargetGeneration::new(4), + AmbiguousResolution::DefinitelyNotAccepted, + ) + .unwrap() + .is_none() + ); + assert_eq!(placement.entry_count(), 0); +} + +#[tokio::test] +async fn explicit_ambiguity_can_be_resolved_as_accepted() { + let placement = placement(1, 32); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + let dispatch = begin(initialization, "cluster-a".to_string(), 6); + let attempt = dispatch.attempt_id(); + + assert!( + dispatch + .finish(PlacementDispatchOutcome::Ambiguous) + .unwrap() + .is_none() + ); + let lease = placement + .resolve_ambiguous( + "session", + attempt, + TargetGeneration::new(6), + AmbiguousResolution::Accepted, + ) + .unwrap() + .unwrap(); + + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-a") + ); + drop(lease); +} + +#[tokio::test] +async fn query_bounds_and_invalidation_are_enforced() { + let placement = placement(1, 4); + let acquire = placement.acquire("one").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = acquire else { + panic!("first acquire should initialize"); + }; + let mut lease = accept(initialization, "cluster-a".to_string(), 1); + + assert_eq!( + query_string(&placement, "one").as_deref(), + Some("cluster-a") + ); + assert!(matches!( + placement.acquire("two").await, + Err(SessionPlacementError::Capacity { max_entries: 1 }) + )); + assert!(matches!( + placement.acquire("12345").await, + Err(SessionPlacementError::KeyTooLong { + actual_bytes: 5, + max_bytes: 4 + }) + )); + + lease.invalidate(); + assert!(placement.query("one").unwrap().is_none()); + assert_eq!(placement.entry_count(), 0); +} + +#[tokio::test] +async fn stale_lease_cannot_remove_a_new_placement() { + let placement = placement(1, 32); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(first) = first else { + panic!("first acquire should initialize"); + }; + let mut invalidating_lease = accept(first, "cluster-a".to_string(), 1); + let second = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Bound { + lease: mut stale_lease, + .. + } = second + else { + panic!("second acquire should use the existing placement"); + }; + + invalidating_lease.invalidate(); + let replacement = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(replacement) = replacement else { + panic!("invalidated placement should be initialized again"); + }; + let replacement_lease = accept(replacement, "cluster-b".to_string(), 2); + + stale_lease.invalidate(); + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-b") + ); + assert_eq!(placement.entry_count(), 1); + drop(replacement_lease); +} + +#[tokio::test(start_paused = true)] +async fn stale_reservation_cannot_overwrite_a_replacement() { + let placement = placement(1, 32); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(first) = first else { + panic!("first acquire should initialize"); + }; + + tokio::time::advance(Duration::from_secs(11)).await; + let replacement = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(replacement) = replacement else { + panic!("expired reservation should be replaced"); + }; + let replacement_lease = accept(replacement, "cluster-b".to_string(), 2); + + assert!(matches!( + first.begin_dispatch("cluster-a".to_string(), TargetGeneration::new(1)), + Err(SessionPlacementError::InitializationChanged + | SessionPlacementError::InitializationCancelled) + )); + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-b") + ); + drop(replacement_lease); +} + +#[tokio::test(start_paused = true)] +async fn expired_reservation_cannot_begin_dispatch() { + let placement = placement(1, 32); + let first = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = first else { + panic!("first acquire should initialize"); + }; + + tokio::time::advance(Duration::from_secs(11)).await; + assert!(matches!( + initialization.begin_dispatch("cluster-a".to_string(), TargetGeneration::new(1)), + Err(SessionPlacementError::InitializationCancelled) + )); + assert!(placement.query("session").unwrap().is_none()); + assert_eq!(placement.entry_count(), 0); +} + +#[tokio::test(start_paused = true)] +async fn active_lease_blocks_expiration_and_release_refreshes_ttl() { + let placement = placement(10, 32); + placement.wait_for_reaper().await; + let acquire = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = acquire else { + panic!("first acquire should initialize"); + }; + let lease = accept(initialization, "cluster-a".to_string(), 1); + + tokio::time::advance(Duration::from_secs(20)).await; + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-a") + ); + drop(lease); + tokio::time::advance(Duration::from_secs(9)).await; + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-a") + ); + tokio::time::advance(Duration::from_secs(1)).await; + tokio::task::yield_now().await; + assert!(placement.query("session").unwrap().is_none()); +} + +#[tokio::test(start_paused = true)] +async fn abandoned_lease_does_not_refresh_ttl() { + let placement = placement(10, 32); + let acquire = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = acquire else { + panic!("first acquire should initialize"); + }; + let mut lease = accept(initialization, "cluster-a".to_string(), 1); + + tokio::time::advance(Duration::from_secs(11)).await; + assert_eq!( + query_string(&placement, "session").as_deref(), + Some("cluster-a") + ); + lease.abandon(); + assert!(placement.query("session").unwrap().is_none()); +} + +#[tokio::test] +async fn capacity_path_reaps_expired_entries() { + let placement = placement(1, 32); + let acquire = placement.acquire("one").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = acquire else { + panic!("first acquire should initialize"); + }; + drop(accept(initialization, "cluster-a".to_string(), 1)); + placement.expire_for_test("one"); + + let acquire = placement.acquire("two").await.unwrap(); + assert!(matches!(&acquire, PlacementAcquire::Initialize(_))); + assert_eq!(placement.entry_count(), 1); + drop(acquire); +} + +#[tokio::test] +async fn expired_bound_entry_is_reinitialized_before_reaper_runs() { + let placement = placement(1, 32); + let acquire = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = acquire else { + panic!("first acquire should initialize"); + }; + drop(accept(initialization, "cluster-a".to_string(), 1)); + placement.expire_for_test("session"); + + let replacement = placement.acquire("session").await.unwrap(); + assert!(matches!(&replacement, PlacementAcquire::Initialize(_))); + assert_eq!(placement.entry_count(), 1); + drop(replacement); + assert_eq!(placement.entry_count(), 0); +} + +#[tokio::test] +async fn target_does_not_need_clone() { + struct NonCloneTarget(&'static str); + + let placement = placement(1, 32); + let acquire = placement.acquire("session").await.unwrap(); + let PlacementAcquire::Initialize(initialization) = acquire else { + panic!("first acquire should initialize"); + }; + let lease = accept(initialization, NonCloneTarget("cluster-a"), 1); + let target = placement.query("session").unwrap().unwrap(); + assert_eq!(target.target().0, "cluster-a"); + drop(lease); +} + +#[tokio::test] +async fn dropping_last_handle_cancels_reaper() { + let placement = placement::(10, 32); + let cancellation = placement.cancellation_token(); + drop(placement); + cancellation.cancelled().await; +} + +#[tokio::test] +async fn reaper_restarts_after_its_task_stops() { + let placement = placement::(10, 32); + placement.wait_for_reaper().await; + placement.stop_reaper_for_test().await; + + assert!(placement.query("session").unwrap().is_none()); + placement.wait_for_reaper().await; +} + +#[tokio::test] +async fn invalid_config_is_rejected() { + for config in [ + SessionPlacementConfig { + idle_ttl: Duration::ZERO, + initialization_timeout: Some(Duration::from_secs(1)), + max_entries: 1, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_nanos(1), + initialization_timeout: Some(Duration::from_secs(1)), + max_entries: 1, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_secs(31_536_001), + initialization_timeout: Some(Duration::from_secs(1)), + max_entries: 1, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_secs(1), + initialization_timeout: Some(Duration::ZERO), + max_entries: 1, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_secs(1), + initialization_timeout: Some(Duration::from_nanos(1)), + max_entries: 1, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_secs(1), + initialization_timeout: Some(Duration::from_secs(31_536_001)), + max_entries: 1, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_secs(1), + initialization_timeout: Some(Duration::from_secs(1)), + max_entries: 0, + max_key_bytes: 1, + }, + SessionPlacementConfig { + idle_ttl: Duration::from_secs(1), + initialization_timeout: Some(Duration::from_secs(1)), + max_entries: 1, + max_key_bytes: 0, + }, + ] { + assert!(matches!( + SessionPlacement::::new(config), + Err(SessionPlacementError::InvalidConfig(_)) + )); + } +} + +#[test] +fn placement_requires_a_tokio_runtime() { + let result = SessionPlacement::::new(SessionPlacementConfig { + idle_ttl: Duration::from_secs(1), + initialization_timeout: Some(Duration::from_secs(1)), + max_entries: 1, + max_key_bytes: 1, + }); + assert!(matches!( + result, + Err(SessionPlacementError::RuntimeUnavailable) + )); +}