diff --git a/components/src/dynamo/kv_dc_relay/README.md b/components/src/dynamo/kv_dc_relay/README.md index 9461a6e887c9..a55803a046ec 100644 --- a/components/src/dynamo/kv_dc_relay/README.md +++ b/components/src/dynamo/kv_dc_relay/README.md @@ -5,14 +5,16 @@ SPDX-License-Identifier: Apache-2.0 # DC KV Relay -The DC KV Relay aggregates exact KV-cache ownership inside one data center and publishes a compact -Cuckoo-filter (CKF) projection for multi-DC routing. It discovers workers through the Dynamo -runtime, consumes their ordered KV events, and supervises one actor-owned producer for each local -routing pool. +The DC KV Relay discovers Dynamo inference pools, consumes their ordered KV events, and supervises +one actor-owned Cuckoo-filter (CKF) producer for each local pool. -A pool is one logical indexer domain in one DC. The domain captures cache compatibility and routing -isolation; the DC identity remains stable across Relay restarts and endpoint replacement. Runtime -endpoints are bindings for a pool rather than part of the CKF publication identity. +A pool is one atomic Dynamo indexer domain in one data center. Its domain captures cache +compatibility and routing isolation. The Relay does not merge KV state from independent endpoints +or deployments into one actor, even when they serve the same canonical model. + +Canonical model names are request-facing bindings. One model can bind to multiple independent +pools, and each pool keeps its own KV stream. LoRA registrations remain attached to the pool of +their backing base model. For each pool, the Relay: @@ -23,26 +25,13 @@ For each pool, the Relay: - Publishes barrier snapshots and sequenced deltas containing absolute packed-bucket images. The full hashes and refcounts stay in the Relay because a CKF fingerprint is lossy, can collide, -and has no owner identity. The global consumer needs only the compact projection required for -cross-DC prefix search. +and has no owner identity. ## Recovery boundaries -Recovery has two stages: - -1. **Worker to Relay:** The Relay shares the normal Dynamo indexer's worker-query recovery path. - Ordered KV events handle live mutations; gaps and source replacement recover exact rank state - before the new source generation becomes active. -2. **Relay to global consumer:** A new or reconnected lane installs a barrier snapshot, then - continues with sequenced absolute bucket-image deltas. A missing delta retires that lane and - requires another snapshot. - -The current component uses the producer lifecycle and exposes local diagnostics. An in-process -adapter exercises the complete producer/consumer protocol today. Non-local gRPC transport and -cross-DC request forwarding are separate global-router integration work. - -For the complete architecture, pool model, consistency contract, and recovery flow, see -[Multi-DC KV Routing and the DC Relay](../../../../docs/fern/components/router/multi-dc-kv-routing.md). +The Relay shares the normal Dynamo indexer's worker-query recovery path. Ordered KV events handle +live mutations; gaps and source replacement recover exact rank state before the new source epoch +becomes active. A fenced pool is withdrawn before its actor stops. ## Usage diff --git a/lib/llm/src/discovery/kv_source_membership.rs b/lib/llm/src/discovery/kv_source_membership.rs index fdb390a67902..34811f3de971 100644 --- a/lib/llm/src/discovery/kv_source_membership.rs +++ b/lib/llm/src/discovery/kv_source_membership.rs @@ -182,6 +182,28 @@ impl KvSourceMembershipView { KvStateEndpointResolution::Ambiguous { .. } => None, } } + + /// Whether runtime fields that determine this source binding still match the view. + /// + /// Metadata-only runtime changes do not require waiting for another source-membership + /// publication. Endpoint mapping and the logical worker/rank set do. + pub(crate) fn matches_binding_inputs( + &self, + runtime_configs: &HashMap, + ) -> bool { + if self.endpoint_resolution + != resolve_kv_state_endpoint(&self.serving_endpoint, runtime_configs.values()) + { + return false; + } + + let mut worker_count = 0usize; + let workers_match = expected_workers(runtime_configs).all(|(worker, _)| { + worker_count = worker_count.saturating_add(1); + self.sources.contains_key(&worker) + }); + workers_match && worker_count == self.sources.len() + } } #[derive(Debug, Clone, PartialEq, Eq, Error)] @@ -356,28 +378,13 @@ where ) -> KvSourceMembershipView { let endpoint_resolution = resolve_kv_state_endpoint(serving_endpoint, runtime_configs.values()); - let workers: Vec<_> = runtime_configs - .iter() - .flat_map(|(&worker_id, config)| { - (0..config.data_parallel_size).filter_map(move |offset| { - config - .data_parallel_start_rank - .checked_add(offset) - .map(|dp_rank| { - ( - WorkerWithDpRank::new(worker_id, dp_rank), - config.enable_local_indexer, - ) - }) - }) - }) - .collect(); + let workers: HashMap<_, _> = expected_workers(runtime_configs).collect(); let sources: HashMap> = match &endpoint_resolution { KvStateEndpointResolution::Resolved(kv_state_endpoint) => workers - .iter() + .keys() .copied() - .map(|(worker, _)| { + .map(|worker| { let key = KvSourceKey::new(kv_state_endpoint.clone(), worker); (worker, self.status(&key)) }) @@ -387,15 +394,12 @@ where endpoints: endpoints.clone(), }; workers - .iter() + .keys() .copied() - .map(|(worker, _)| (worker, KvSourceStatus::Ambiguous(ambiguity.clone()))) + .map(|worker| (worker, KvSourceStatus::Ambiguous(ambiguity.clone()))) .collect() } }; - let recovery_expected = workers - .into_iter() - .collect::>(); let kv_event_publishing_enabled = runtime_configs .iter() .map(|(&worker_id, config)| (worker_id, config.kv_event_publishing_enabled)) @@ -405,13 +409,31 @@ where serving_endpoint: serving_endpoint.clone(), endpoint_resolution, lifecycle_generations: sources.keys().map(|worker| (*worker, 0)).collect(), - recovery_expected, + recovery_expected: workers, kv_event_publishing_enabled, sources, } } } +fn expected_workers( + runtime_configs: &HashMap, +) -> impl Iterator + '_ { + runtime_configs.iter().flat_map(|(&worker_id, config)| { + (0..config.data_parallel_size).filter_map(move |offset| { + config + .data_parallel_start_rank + .checked_add(offset) + .map(|dp_rank| { + ( + WorkerWithDpRank::new(worker_id, dp_rank), + config.enable_local_indexer, + ) + }) + }) + }) +} + /// Resolve the effective KV-state endpoint advertised by active base runtime configs. /// /// An omitted mapping and an explicit mapping to `serving_endpoint` are equal after fallback. @@ -552,6 +574,35 @@ mod tests { } } + #[test] + fn binding_inputs_ignore_metadata_only_runtime_changes() { + let serving = endpoint("generate"); + let kv_endpoint = endpoint("kv-events"); + let original = HashMap::from([( + 7, + ModelRuntimeConfig { + context_length: Some(4096), + data_parallel_start_rank: 2, + data_parallel_size: 2, + kv_state_endpoint: Some(kv_endpoint.clone()), + ..Default::default() + }, + )]); + let view = KvSourceMembership::::new().view(&serving, &original); + + let mut metadata_only = original.clone(); + metadata_only.get_mut(&7).unwrap().context_length = Some(8192); + assert!(view.matches_binding_inputs(&metadata_only)); + + let mut remapped = metadata_only.clone(); + remapped.get_mut(&7).unwrap().kv_state_endpoint = Some(endpoint("other-kv-events")); + assert!(!view.matches_binding_inputs(&remapped)); + + let mut resized = metadata_only; + resized.get_mut(&7).unwrap().data_parallel_size = 3; + assert!(!view.matches_binding_inputs(&resized)); + } + #[test] fn overlapping_random_incarnations_are_ambiguous_until_one_remains() { let kv_endpoint = endpoint("kv-events"); diff --git a/lib/llm/src/kv_dc_relay.rs b/lib/llm/src/kv_dc_relay.rs index 20fa3fab70de..8c479032b4d3 100644 --- a/lib/llm/src/kv_dc_relay.rs +++ b/lib/llm/src/kv_dc_relay.rs @@ -6,12 +6,13 @@ mod actor; mod discovery; mod host; +mod identity; +mod pool_registry; mod resolution; pub use host::{ DEFAULT_EXPECTED_UNIQUE_BLOCKS, KvDcRelay, KvDcRelayConfig, KvDcRelayError, KvDcRelayHealth, }; - #[cfg(feature = "ckf-diagnostics")] pub use host::{ KvDcRelayActorStats, KvDcRelayAggregationStats, KvDcRelayCacheDomainStats, @@ -19,3 +20,8 @@ pub use host::{ KvDcRelayMemberStats, KvDcRelayMemoryStats, KvDcRelayPublicationStats, KvDcRelayRecoveryStats, KvDcRelayStats, }; +pub use identity::{ + CanonicalModelId, CanonicalModelIdError, CanonicalModelRegistration, DcPoolCatalog, + DcPoolDescriptor, DcRelayIdentity, ModelAlias, ModelAliasError, ModelTarget, + PoolIdentitySources, +}; diff --git a/lib/llm/src/kv_dc_relay/actor.rs b/lib/llm/src/kv_dc_relay/actor.rs index 4c8e4a13c122..57404dcb2b39 100644 --- a/lib/llm/src/kv_dc_relay/actor.rs +++ b/lib/llm/src/kv_dc_relay/actor.rs @@ -12,13 +12,15 @@ use std::sync::atomic::{AtomicU64, Ordering}; #[cfg(feature = "ckf-diagnostics")] use std::time::Instant; +#[cfg(test)] +use dynamo_kv_router::indexer::cuckoo::CkfConfig; #[cfg(any(test, feature = "ckf-diagnostics"))] use dynamo_kv_router::indexer::cuckoo::DcCkfStats; #[cfg(feature = "ckf-diagnostics")] use dynamo_kv_router::indexer::cuckoo::PublisherEmitOutcome; use dynamo_kv_router::indexer::cuckoo::{ - CkfConfig, CkfFailureAction, CkfFailureDisposition, CkfFailurePoint, DcCkfDelta, - DcCkfDeltaSink, DcCkfPublisher, DcCkfSnapshot, DcCkfState, LaneLease, ProducerIdentity, + CkfFailureAction, CkfFailureDisposition, CkfFailurePoint, DcCkfDelta, DcCkfDeltaSink, + DcCkfPublisher, DcCkfSnapshot, DcCkfState, LaneLease, ProducerIdentity, }; use dynamo_kv_router::protocols::{ DpRank, ExternalSequenceBlockHash, KvCacheEventData, KvCacheEventError, RouterEvent, @@ -32,12 +34,11 @@ use tokio_util::sync::CancellationToken; use crate::kv_router::indexer::{RecoveryResetReason, RecoveryTarget, SourceEpoch}; use super::host::KvDcRelayError; -use super::resolution::PoolBinding; const DEFAULT_MAILBOX_CAPACITY: usize = 256; const DEFAULT_PENDING_BLOCK_PERMITS: usize = 65_536; const DEFAULT_PUBLICATION_CAPACITY: usize = 64; -const DEFAULT_FAULT_CAPACITY: usize = 16; +pub(super) const DEFAULT_FAULT_CAPACITY: usize = 16; #[cfg(test)] const DEFAULT_PUBLICATION_DELAY: Duration = Duration::from_millis(1); const RECOVERY_REBUILD_BATCH_WINDOW: Duration = Duration::from_millis(5); @@ -253,11 +254,23 @@ fn actor_fault_category(disposition: CkfFailureDisposition) -> ActorFaultCategor } } +async fn send_actor_fault( + sender: &mpsc::Sender, + fence: &CancellationToken, + fault: ActorFault, +) -> bool { + tokio::select! { + biased; + _ = fence.cancelled() => false, + result = sender.send(fault) => result.is_ok(), + } +} + #[derive(Debug, Clone)] pub(super) struct StreamScope { - pub(super) process_incarnation: u64, + pub(super) relay_incarnation: u64, pub(super) layout_generation: u64, - pub(super) pool_binding: PoolBinding, + pub(super) pool_id: dynamo_kv_router::identity::PoolId, } #[derive(Debug, Clone)] @@ -276,12 +289,12 @@ impl DcCkfDeltaSink for BroadcastDeltaSink { #[derive(Debug, Clone)] pub(crate) struct KvDcRelayHandle { sender: mpsc::Sender, + identity: ProducerIdentity, payload_permits: Arc, fence: CancellationToken, stopped: CancellationToken, #[cfg(feature = "ckf-diagnostics")] pub(super) diagnostics: ActorDiagnosticsHandle, - pub(super) scope: StreamScope, } impl KvDcRelayHandle { @@ -298,6 +311,7 @@ impl KvDcRelayHandle { ) } + #[cfg(test)] pub(super) fn spawn_with_publication_delay( config: CkfConfig, scope: StreamScope, @@ -311,6 +325,19 @@ impl KvDcRelayHandle { ) } + pub(super) fn spawn_with_state_and_publication_delay( + state: DcCkfState, + scope: StreamScope, + publication_delay: Duration, + ) -> (Self, mpsc::Receiver) { + Self::spawn_with_state_capacity_and_delay( + state, + scope, + DEFAULT_MAILBOX_CAPACITY, + publication_delay, + ) + } + #[cfg(test)] fn spawn_with_capacity( config: CkfConfig, @@ -320,6 +347,7 @@ impl KvDcRelayHandle { Self::spawn_with_capacity_and_delay(config, scope, capacity, DEFAULT_PUBLICATION_DELAY) } + #[cfg(test)] fn spawn_with_capacity_and_delay( config: CkfConfig, scope: StreamScope, @@ -327,11 +355,25 @@ impl KvDcRelayHandle { publication_delay: Duration, ) -> Result<(Self, mpsc::Receiver), KvDcRelayError> { let state = DcCkfState::new(config)?; + Ok(Self::spawn_with_state_capacity_and_delay( + state, + scope, + capacity, + publication_delay, + )) + } + + fn spawn_with_state_capacity_and_delay( + state: DcCkfState, + scope: StreamScope, + capacity: usize, + publication_delay: Duration, + ) -> (Self, mpsc::Receiver) { let (sender, receiver) = mpsc::channel(capacity); let (publication_tx, _) = broadcast::channel(DEFAULT_PUBLICATION_CAPACITY); let identity = ProducerIdentity::new( - scope.pool_binding.pool_id(), - scope.process_incarnation, + scope.pool_id, + scope.relay_incarnation, scope.layout_generation, state.format(), ); @@ -356,18 +398,22 @@ impl KvDcRelayHandle { fence.clone(), stopped.clone(), )); - Ok(( + ( Self { sender, + identity, payload_permits: Arc::new(Semaphore::new(DEFAULT_PENDING_BLOCK_PERMITS)), fence, stopped, #[cfg(feature = "ckf-diagnostics")] diagnostics, - scope, }, fault_rx, - )) + ) + } + + pub(super) const fn identity(&self) -> ProducerIdentity { + self.identity } async fn submit( @@ -930,8 +976,10 @@ async fn run_actor( source_epoch.get() ); diagnostics.record_error(&message); - if fault_tx - .send(ActorFault { + if !send_actor_fault( + &fault_tx, + &fence, + ActorFault { worker_id, dp_rank, source_epoch, @@ -939,10 +987,11 @@ async fn run_actor( category: actor_fault_category(disposition), disposition, message, - }) - .await - .is_err() + }, + ) + .await { + discard_tail = fence.is_cancelled(); break; } diagnostics.finish_command(); @@ -1023,8 +1072,10 @@ async fn run_actor( let message = error.to_string(); let category = actor_fault_category(disposition); diagnostics.record_error(&message); - if fault_tx - .send(ActorFault { + if !send_actor_fault( + &fault_tx, + &fence, + ActorFault { worker_id, dp_rank, source_epoch, @@ -1032,10 +1083,11 @@ async fn run_actor( category, disposition, message, - }) - .await - .is_err() + }, + ) + .await { + discard_tail = fence.is_cancelled(); break; } } @@ -1314,27 +1366,19 @@ mod tests { use dynamo_kv_router::protocols::{ KvCacheEvent, KvCacheStoreData, KvCacheStoredBlockData, LocalBlockHash, }; - use dynamo_runtime::protocols::EndpointId; use super::*; - use crate::kv_dc_relay::resolution::EndpointLocator; - fn scope(name: &str) -> StreamScope { - let endpoint = format!("ns.worker.{name}"); - let endpoint_id = EndpointId::from(endpoint.as_str()); + fn scope(_name: &str) -> StreamScope { let dc_id = DcId::new(2); let domain = IndexerDomainId::new( CacheSemanticsId::new([1; 16], IdentitySource::Explicit), RoutingScopeId::new([3; 16], IdentitySource::Explicit), ); StreamScope { - process_incarnation: 1, + relay_incarnation: 1, layout_generation: 1, - pool_binding: PoolBinding::new( - PoolId::new(domain, dc_id), - EndpointLocator::new(dc_id, endpoint_id), - None, - ), + pool_id: PoolId::new(domain, dc_id), } } @@ -1399,10 +1443,7 @@ mod tests { snapshot.buckets.len(), snapshot.identity.format().bucket_count() ); - assert_eq!( - actor_health(&handle).mailbox_capacity, - DEFAULT_MAILBOX_CAPACITY - ); + assert_eq!(handle.mailbox_capacity(), DEFAULT_MAILBOX_CAPACITY); assert!( handle .diagnostics @@ -1597,6 +1638,37 @@ mod tests { )); } + #[tokio::test] + async fn producer_fence_interrupts_a_full_fault_channel() { + let worker = WorkerWithDpRank::new(1, 0); + let (handle, faults) = + KvDcRelayHandle::spawn(CkfConfig::new(32), scope("fault-backpressure")).unwrap(); + handle + .admit_event(SourceEpoch::new(0), stored(worker, 1, &[1])) + .await + .unwrap(); + handle.flush().await.unwrap(); + + for event_id in 2..=(DEFAULT_FAULT_CAPACITY as u64 + 2) { + handle + .admit_event(SourceEpoch::new(1), stored(worker, event_id, &[event_id])) + .await + .unwrap(); + } + tokio::time::timeout(Duration::from_secs(1), async { + while faults.len() != DEFAULT_FAULT_CAPACITY || handle.mailbox_depth() != 0 { + tokio::task::yield_now().await; + } + }) + .await + .expect("actor must block after filling the fault channel"); + + tokio::time::timeout(Duration::from_secs(1), handle.fence()) + .await + .expect("fence must interrupt a blocked fault send") + .unwrap(); + } + #[tokio::test] async fn cadence_advances_on_duplicate_events_without_acknowledging_mutation() { let worker = WorkerWithDpRank::new(1, 0); diff --git a/lib/llm/src/kv_dc_relay/discovery.rs b/lib/llm/src/kv_dc_relay/discovery.rs index 6ad9601b75e7..b4fe15cee1ec 100644 --- a/lib/llm/src/kv_dc_relay/discovery.rs +++ b/lib/llm/src/kv_dc_relay/discovery.rs @@ -2,6 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 use std::collections::{HashMap, HashSet}; +use std::pin::Pin; use std::sync::Arc; use std::time::Duration; @@ -11,11 +12,12 @@ use dynamo_runtime::discovery::{ ModelCardInstanceId, }; use dynamo_runtime::protocols::EndpointId; -use futures::StreamExt; +use futures::{Stream, StreamExt, future::try_join_all}; use tokio::sync::watch; use tokio::task::JoinHandle; use tokio_util::sync::CancellationToken; +use super::identity::{CanonicalModelId, CanonicalModelRegistration, ModelAlias, ModelTarget}; use super::resolution::{ResolvedIndexerDomain, resolve_indexer_domain}; use crate::local_model::runtime_config::ModelRuntimeConfig; use crate::model_card::ModelDeploymentCard; @@ -25,46 +27,242 @@ const KV_EVENT_HASH_FORMAT_VERSION: u16 = 1; pub(crate) type KvCacheDomainKey = ResolvedIndexerDomain; +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct KvDcRelayDiscoveryConfig { + pub namespaces: Vec, + pub endpoint_prefixes: Vec, + pub watch_all: bool, +} + +impl KvDcRelayDiscoveryConfig { + pub fn validate(&self) -> anyhow::Result<()> { + anyhow::ensure!( + self.watch_all || !self.namespaces.is_empty(), + "KV DC Relay requires at least one discovery namespace or explicit watch_all" + ); + anyhow::ensure!( + !self.watch_all || self.namespaces.is_empty(), + "KV DC Relay watch_all cannot be combined with explicit discovery namespaces" + ); + + let mut unique_namespaces = HashSet::new(); + for namespace in &self.namespaces { + anyhow::ensure!( + !namespace.trim().is_empty(), + "KV DC Relay discovery namespaces must not be empty" + ); + anyhow::ensure!( + namespace.trim() == namespace, + "KV DC Relay discovery namespaces must not contain surrounding whitespace" + ); + anyhow::ensure!( + unique_namespaces.insert(namespace), + "duplicate KV DC Relay discovery namespace: {namespace}" + ); + } + + let mut unique_prefixes = HashSet::new(); + for prefix in &self.endpoint_prefixes { + anyhow::ensure!( + !prefix.trim().is_empty(), + "KV DC Relay endpoint prefixes must not be empty" + ); + anyhow::ensure!( + prefix.trim() == prefix, + "KV DC Relay endpoint prefixes must not contain surrounding whitespace" + ); + anyhow::ensure!( + unique_prefixes.insert(prefix), + "duplicate KV DC Relay endpoint prefix: {prefix}" + ); + anyhow::ensure!( + self.watch_all + || self.namespaces.iter().any(|namespace| { + prefix == namespace + || prefix + .strip_prefix(namespace) + .is_some_and(|suffix| suffix.starts_with('.')) + }), + "KV DC Relay endpoint prefix {prefix} is outside the configured namespaces" + ); + } + Ok(()) + } + + fn queries(&self) -> Vec { + if self.watch_all { + vec![DiscoveryQuery::AllModels] + } else { + self.namespaces + .iter() + .map(|namespace| DiscoveryQuery::NamespacedModels { + namespace: namespace.clone(), + }) + .collect() + } + } + + fn filter(&self) -> DcDiscoveryFilter { + DcDiscoveryFilter { + endpoint_prefixes: self.endpoint_prefixes.clone(), + } + } +} + #[derive(Debug, Clone, Default)] pub(crate) struct DcDiscoveryFilter { - pub(crate) namespace: Option, - pub(crate) endpoint_prefix: Option, + endpoint_prefixes: Vec, } impl DcDiscoveryFilter { fn matches(&self, endpoint: &EndpointId) -> bool { - if self - .namespace - .as_ref() - .is_some_and(|namespace| endpoint.namespace != *namespace) - { - return false; + if self.endpoint_prefixes.is_empty() { + return true; } - self.endpoint_prefix.as_ref().is_none_or(|prefix| { - format!( - "{}.{}.{}", - endpoint.namespace, endpoint.component, endpoint.name - ) - .starts_with(prefix) - }) + self.endpoint_prefixes + .iter() + .any(|prefix| endpoint_matches_prefix(endpoint, prefix)) } } +fn endpoint_matches_prefix(endpoint: &EndpointId, prefix: &str) -> bool { + let mut parts = prefix.split('.'); + for expected in [ + endpoint.namespace.as_str(), + endpoint.component.as_str(), + endpoint.name.as_str(), + ] { + match parts.next() { + None => return true, + Some(actual) if actual == expected => {} + Some(_) => return false, + } + } + parts.next().is_none() +} + +/// A structural inconsistency that makes the endpoint unsafe to materialize. +/// +/// Invalid aliases and orphan adapter cards are soft discovery errors: they are omitted and +/// logged, but do not create a conflict for an otherwise valid base endpoint. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum MaterializationConflict { + Endpoint { + endpoint: EndpointId, + reason: String, + }, + Card { + card: ModelCardInstanceId, + reason: String, + }, +} + #[derive(Debug, Clone, PartialEq)] pub(crate) struct EndpointMembership { pub(crate) endpoint: EndpointId, pub(crate) generation: u64, pub(crate) domain: Option, - pub(crate) compatibility_conflict: bool, + pub(crate) registrations: Vec, pub(crate) models: Vec, pub(crate) aliases: Vec, pub(crate) roles: Vec, pub(crate) runtime_configs: HashMap, + pub(crate) conflicts: Vec, +} + +impl EndpointMembership { + pub(crate) fn is_materializable(&self) -> bool { + self.domain.is_some() && !self.registrations.is_empty() && self.conflicts.is_empty() + } } #[derive(Debug, Clone, Default, PartialEq)] pub(crate) struct DcMembershipView { - pub(crate) endpoints: HashMap, + pub(crate) endpoints: Arc>, +} + +#[derive(Debug, Clone)] +struct ProjectedBaseCard<'a> { + id: &'a ModelCardInstanceId, + card: &'a ModelDeploymentCard, + domain: KvCacheDomainKey, + model: Option, + aliases: Vec, +} + +#[derive(Debug)] +struct StoredModelCard { + card: ModelDeploymentCard, + serialized: serde_json::Value, +} + +impl PartialEq for StoredModelCard { + fn eq(&self, other: &Self) -> bool { + self.serialized == other.serialized + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)] +struct BindingIdentity { + model: CanonicalModelId, + target: ModelTarget, +} + +#[derive(Debug, Clone)] +struct RegistrationClaim { + binding: BindingIdentity, + aliases: Vec, +} + +struct EndpointMembershipBuilder { + domain: Option, + claims: Vec, + runtime_configs: HashMap, + roles: HashSet, + conflicts: Vec, +} + +struct ProjectionDiagnostics<'a> { + invalid_models: HashSet, + invalid_aliases: HashSet<(ModelCardInstanceId, String)>, + orphan_adapters: HashSet, + warned_invalid_models: &'a HashSet, + warned_invalid_aliases: &'a HashSet<(ModelCardInstanceId, String)>, + warned_orphan_adapters: &'a HashSet, +} + +impl<'a> ProjectionDiagnostics<'a> { + fn new( + warned_invalid_models: &'a HashSet, + warned_invalid_aliases: &'a HashSet<(ModelCardInstanceId, String)>, + warned_orphan_adapters: &'a HashSet, + ) -> Self { + Self { + invalid_models: HashSet::new(), + invalid_aliases: HashSet::new(), + orphan_adapters: HashSet::new(), + warned_invalid_models, + warned_invalid_aliases, + warned_orphan_adapters, + } + } + + fn valid_aliases( + &mut self, + id: &ModelCardInstanceId, + card: &ModelDeploymentCard, + model: &CanonicalModelId, + endpoint: &EndpointId, + ) -> Vec { + valid_aliases( + id, + card, + model, + endpoint, + &mut self.invalid_aliases, + self.warned_invalid_aliases, + ) + } } pub(crate) struct DcMembershipWatch { @@ -76,17 +274,20 @@ pub(crate) struct DcMembershipWatch { impl DcMembershipWatch { pub(crate) async fn start( discovery: Arc, - filter: DcDiscoveryFilter, + config: KvDcRelayDiscoveryConfig, parent_cancel: CancellationToken, ) -> anyhow::Result { - let initial = discovery.list(DiscoveryQuery::AllModels).await?; + config.validate()?; + let queries = config.queries(); + let filter = config.filter(); + let initial = list_queries(&discovery, &queries).await?; let mut state = MembershipState::default(); state.replace_all(initial, &filter); let (sender, receiver) = watch::channel(state.view(&filter)); let cancel = parent_cancel.child_token(); let task_cancel = cancel.clone(); let task = tokio::spawn(async move { - run_membership_watch(discovery, filter, state, sender, task_cancel).await; + run_membership_watch(discovery, queries, filter, state, sender, task_cancel).await; }); Ok(Self { receiver, @@ -109,15 +310,38 @@ impl DcMembershipWatch { } } -#[derive(Default)] struct MembershipState { - cards: HashMap, - endpoint_generations: HashMap, - previous: HashMap, + cards: HashMap, + next_membership_generation: u64, + previous: Arc>, + warned_invalid_models: HashSet, + warned_invalid_aliases: HashSet<(ModelCardInstanceId, String)>, + warned_orphan_adapters: HashSet, + #[cfg(test)] + projection_count: usize, +} + +impl Default for MembershipState { + fn default() -> Self { + Self { + cards: HashMap::new(), + next_membership_generation: 1, + previous: Arc::default(), + warned_invalid_models: HashSet::new(), + warned_invalid_aliases: HashSet::new(), + warned_orphan_adapters: HashSet::new(), + #[cfg(test)] + projection_count: 0, + } + } } impl MembershipState { - fn replace_all(&mut self, instances: Vec, filter: &DcDiscoveryFilter) { + fn replace_all( + &mut self, + instances: Vec, + filter: &DcDiscoveryFilter, + ) -> bool { let mut next = HashMap::new(); for instance in instances { let Some((id, card)) = decode_card(instance) else { @@ -127,114 +351,354 @@ impl MembershipState { next.insert(id, card); } } + if self.cards == next { + return false; + } self.cards = next; + true } - fn apply(&mut self, event: DiscoveryEvent, filter: &DcDiscoveryFilter) { + fn apply(&mut self, event: DiscoveryEvent, filter: &DcDiscoveryFilter) -> bool { match event { DiscoveryEvent::Added(instance) => { let Some((id, card)) = decode_card(instance) else { - return; + return false; }; if filter.matches(&endpoint_id(&id)) { + if self.cards.get(&id) == Some(&card) { + return false; + } self.cards.insert(id, card); + return true; } + false } DiscoveryEvent::Removed(DiscoveryInstanceId::Model(id)) => { - self.cards.remove(&id); + self.cards.remove(&id).is_some() } - DiscoveryEvent::Removed(_) => {} + DiscoveryEvent::Removed(_) => false, } } fn view(&mut self, filter: &DcDiscoveryFilter) -> DcMembershipView { + #[cfg(test)] + { + self.projection_count = self.projection_count.saturating_add(1); + } + let mut diagnostics = ProjectionDiagnostics::new( + &self.warned_invalid_models, + &self.warned_invalid_aliases, + &self.warned_orphan_adapters, + ); let mut grouped: HashMap> = HashMap::new(); - for (id, card) in &self.cards { + for (id, stored) in &self.cards { let endpoint = endpoint_id(id); if filter.matches(&endpoint) { - grouped.entry(endpoint).or_default().push((id, card)); + grouped + .entry(endpoint) + .or_default() + .push((id, &stored.card)); } } - let mut endpoints = HashMap::new(); + let mut builders = HashMap::::new(); for (endpoint, cards) in grouped { - let mut domains = HashSet::new(); - let mut models = HashSet::new(); - let mut aliases = HashSet::new(); - let mut roles = HashSet::new(); - let mut runtime_configs = HashMap::new(); - + let mut base_cards = Vec::new(); + let mut adapter_cards = Vec::new(); for (id, card) in cards { - models.insert(card.name().to_string()); - aliases.extend(card.aliases.iter().cloned()); - if let Some(role) = card.worker_type { - roles.insert(format!("{role:?}").to_lowercase()); - } if id.model_suffix.is_some() || card.lora.is_some() { + adapter_cards.push((id, card)); continue; } - domains.insert(resolve_indexer_domain( + let domain = resolve_indexer_domain(card, &endpoint, KV_EVENT_HASH_FORMAT_VERSION); + let model = match CanonicalModelId::new(card.name().to_string()) { + Ok(model) => Some(model), + Err(error) => { + diagnostics.invalid_models.insert(id.clone()); + if !diagnostics.warned_invalid_models.contains(id) { + tracing::warn!( + endpoint = %endpoint, + model = card.name(), + %error, + "Ignoring model card with invalid canonical model identity" + ); + } + None + } + }; + let aliases = model + .as_ref() + .map(|model| diagnostics.valid_aliases(id, card, model, &endpoint)) + .unwrap_or_default(); + base_cards.push(ProjectedBaseCard { + id, card, - &endpoint, - KV_EVENT_HASH_FORMAT_VERSION, - )); - runtime_configs.insert(id.instance_id, card.runtime_config.clone()); + domain, + model, + aliases, + }); } - let compatibility_conflict = domains.len() > 1; - let domain = (domains.len() == 1) - .then(|| domains.into_iter().next()) + let endpoint_domains: HashSet<_> = base_cards + .iter() + .map(|projection| projection.domain.clone()) + .collect(); + let domain_count = endpoint_domains.len(); + let domain = (domain_count == 1) + .then(|| endpoint_domains.into_iter().next()) .flatten(); + let mut builder = EndpointMembershipBuilder { + domain, + claims: Vec::new(), + runtime_configs: HashMap::new(), + roles: HashSet::new(), + conflicts: Vec::new(), + }; + if domain_count > 1 { + builder.conflicts.push(MaterializationConflict::Endpoint { + endpoint: endpoint.clone(), + reason: "endpoint resolves to multiple indexer domains".to_string(), + }); + } + + let base_models: HashSet<_> = base_cards + .iter() + .filter_map(|projection| projection.model.clone()) + .collect(); + if base_models.len() > 1 { + builder.conflicts.push(MaterializationConflict::Endpoint { + endpoint: endpoint.clone(), + reason: "endpoint resolves to multiple canonical base models".to_string(), + }); + } + + let mut base_by_worker = HashMap::with_capacity(base_cards.len()); + for projection in &base_cards { + let worker_id = projection.id.instance_id; + debug_assert!(base_by_worker.insert(worker_id, projection).is_none()); + let Some(model) = projection.model.clone() else { + builder.conflicts.push(MaterializationConflict::Card { + card: projection.id.clone(), + reason: "invalid canonical model identity".to_string(), + }); + continue; + }; + builder + .runtime_configs + .insert(worker_id, projection.card.runtime_config.clone()); + if let Some(worker_type) = projection.card.worker_type { + builder.roles.insert(worker_type.as_str().to_string()); + } + builder.claims.push(RegistrationClaim { + binding: BindingIdentity { + model: model.clone(), + target: ModelTarget::Base { base_model: model }, + }, + aliases: projection.aliases.clone(), + }); + } + + for (id, card) in adapter_cards { + let worker_id = id.instance_id; + let Some(base) = base_by_worker.get(&worker_id).copied() else { + diagnostics.orphan_adapters.insert(id.clone()); + if !diagnostics.warned_orphan_adapters.contains(id) { + tracing::warn!( + endpoint = %endpoint, + worker_id, + card = %id.to_path(), + "Ignoring adapter card without a backing base model card" + ); + } + continue; + }; + let Some(base_model) = base.model.clone() else { + builder.conflicts.push(MaterializationConflict::Card { + card: id.clone(), + reason: "adapter's backing base model identity is invalid".to_string(), + }); + continue; + }; + let adapter_name = card + .lora + .as_ref() + .map(|lora| lora.name.as_str()) + .or(id.model_suffix.as_deref()); + let Some(adapter_name) = adapter_name else { + builder.conflicts.push(MaterializationConflict::Card { + card: id.clone(), + reason: "adapter card has no adapter identity".to_string(), + }); + continue; + }; + let adapter = match CanonicalModelId::new(adapter_name.to_string()) { + Ok(adapter) => adapter, + Err(error) => { + diagnostics.invalid_models.insert(id.clone()); + if !diagnostics.warned_invalid_models.contains(id) { + tracing::warn!( + endpoint = %endpoint, + model = adapter_name, + %error, + "Ignoring adapter card with invalid canonical model identity" + ); + } + builder.conflicts.push(MaterializationConflict::Card { + card: id.clone(), + reason: "invalid adapter model identity".to_string(), + }); + continue; + } + }; + let aliases = diagnostics.valid_aliases(id, card, &adapter, &endpoint); + let target = ModelTarget::Lora { + base_model, + adapter: adapter.clone(), + }; + builder.claims.push(RegistrationClaim { + binding: BindingIdentity { + model: adapter, + target, + }, + aliases, + }); + } + builders.insert(endpoint, builder); + } + + let mut endpoints = HashMap::new(); + for (endpoint, mut builder) in builders { + let mut lookup_owners = HashMap::>::new(); + for claim in &builder.claims { + lookup_owners + .entry(claim.binding.model.as_str().to_string()) + .or_default() + .insert(claim.binding.clone()); + for alias in &claim.aliases { + lookup_owners + .entry(alias.as_str().to_string()) + .or_default() + .insert(claim.binding.clone()); + } + } + for name in lookup_owners + .into_iter() + .filter_map(|(name, owners)| (owners.len() > 1).then_some(name)) + { + builder.conflicts.push(MaterializationConflict::Endpoint { + endpoint: endpoint.clone(), + reason: format!( + "request-facing name {name} resolves to conflicting targets within the endpoint pool" + ), + }); + } + + let mut grouped_claims = HashMap::>::new(); + for claim in builder.claims { + let aliases = grouped_claims.entry(claim.binding).or_default(); + aliases.extend(claim.aliases); + } + let mut registrations: Vec<_> = grouped_claims + .into_iter() + .map(|(binding, aliases)| { + CanonicalModelRegistration::with_target( + binding.model, + binding.target, + aliases.into_iter().collect(), + ) + }) + .collect(); + registrations.sort_unstable(); + let models = registrations + .iter() + .map(|registration| registration.model().as_str().to_string()) + .collect::>(); + let aliases = registrations + .iter() + .flat_map(|registration| registration.aliases()) + .map(|alias| alias.as_str().to_string()) + .collect::>(); + builder + .conflicts + .sort_by(|left, right| format!("{left:?}").cmp(&format!("{right:?}"))); + builder.conflicts.dedup(); let mut candidate = EndpointMembership { endpoint: endpoint.clone(), - generation: 0, - domain, - compatibility_conflict, + generation: self + .previous + .get(&endpoint) + .map_or(0, |previous| previous.generation), + domain: builder.domain, + registrations, models: sorted(models), aliases: sorted(aliases), - roles: sorted(roles), - runtime_configs, + roles: sorted(builder.roles), + runtime_configs: builder.runtime_configs, + conflicts: builder.conflicts, }; - let changed = self - .previous - .get(&endpoint) - .is_none_or(|previous| !same_membership(previous, &candidate)); - let generation = self - .endpoint_generations - .entry(endpoint.clone()) - .or_default(); - if changed { - *generation = generation.saturating_add(1); + if self.previous.get(&endpoint) != Some(&candidate) { + candidate.generation = self.next_membership_generation; + self.next_membership_generation = self.next_membership_generation.saturating_add(1); } - candidate.generation = *generation; endpoints.insert(endpoint, candidate); } - self.endpoint_generations - .retain(|endpoint, _| endpoints.contains_key(endpoint)); + let ProjectionDiagnostics { + invalid_models, + invalid_aliases, + orphan_adapters, + .. + } = diagnostics; + self.warned_invalid_models = invalid_models; + self.warned_invalid_aliases = invalid_aliases; + self.warned_orphan_adapters = orphan_adapters; + let endpoints = Arc::new(endpoints); self.previous = endpoints.clone(); DcMembershipView { endpoints } } } +#[cfg(test)] +pub(super) fn project_instances_for_test(instances: Vec) -> DcMembershipView { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + state.replace_all(instances, &filter); + state.view(&filter) +} + async fn run_membership_watch( discovery: Arc, + queries: Vec, filter: DcDiscoveryFilter, mut state: MembershipState, sender: watch::Sender, cancel: CancellationToken, ) { let mut retry_delay = Duration::from_millis(100); + let mut watch_failures = 0u64; + let mut reconcile_failures = 0u64; loop { let stream_cancel = cancel.child_token(); - let stream = discovery - .list_and_watch(DiscoveryQuery::AllModels, Some(stream_cancel.clone())) - .await; - let mut stream = match stream { + let streams = open_query_streams(&discovery, &queries, &stream_cancel).await; + let mut stream = match streams { Ok(stream) => stream, Err(error) => { - tracing::error!(%error, "Failed to watch DC-wide model-card membership"); + stream_cancel.cancel(); + watch_failures = watch_failures.saturating_add(1); + if watch_failures == 1 { + tracing::error!( + %error, + query_count = queries.len(), + "Failed to watch scoped KV DC Relay model-card membership" + ); + } else { + tracing::debug!( + %error, watch_failures, query_count = queries.len(), + retry_ms = retry_delay.as_millis(), + "Scoped KV DC Relay model-card watch retry failed" + ); + } if !retry_or_cancel(retry_delay, &cancel).await { return; } @@ -242,7 +706,6 @@ async fn run_membership_watch( continue; } }; - retry_delay = Duration::from_millis(100); let mut reconcile = tokio::time::interval(RECONCILE_INTERVAL); reconcile.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); @@ -250,22 +713,58 @@ async fn run_membership_watch( tokio::select! { _ = cancel.cancelled() => return, event = stream.next() => match event { - Some(Ok(event)) => { - state.apply(event, &filter); - sender.send_replace(state.view(&filter)); + Some(Ok(Some(event))) => { + watch_failures = 0; + retry_delay = Duration::from_millis(100); + if state.apply(event, &filter) { + publish_membership_if_changed(&sender, state.view(&filter)); + } } Some(Err(error)) => { - tracing::error!(%error, "DC-wide model-card discovery stream failed; rebinding"); + watch_failures = watch_failures.saturating_add(1); + if watch_failures == 1 { + tracing::error!(%error, "Scoped KV DC Relay model-card discovery stream failed; rebinding"); + } else { + tracing::debug!( + %error, watch_failures, retry_ms = retry_delay.as_millis(), + "Scoped KV DC Relay model-card discovery stream failed again; rebinding" + ); + } + break; + } + Some(Ok(None)) | None => { + watch_failures = watch_failures.saturating_add(1); + if watch_failures == 1 { + tracing::error!("Scoped KV DC Relay model-card discovery stream closed; rebinding"); + } else { + tracing::debug!( + watch_failures, retry_ms = retry_delay.as_millis(), + "Scoped KV DC Relay model-card discovery stream closed again; rebinding" + ); + } break; } - None => break, }, - _ = reconcile.tick() => match discovery.list(DiscoveryQuery::AllModels).await { + _ = reconcile.tick() => match list_queries(&discovery, &queries).await { Ok(instances) => { - state.replace_all(instances, &filter); - sender.send_replace(state.view(&filter)); + watch_failures = 0; + reconcile_failures = 0; + retry_delay = Duration::from_millis(100); + if state.replace_all(instances, &filter) { + publish_membership_if_changed(&sender, state.view(&filter)); + } + } + Err(error) => { + reconcile_failures = reconcile_failures.saturating_add(1); + if reconcile_failures == 1 { + tracing::warn!(%error, "Failed periodic KV DC Relay membership reconciliation"); + } else { + tracing::debug!( + %error, reconcile_failures, + "Periodic KV DC Relay membership reconciliation failed again" + ); + } } - Err(error) => tracing::warn!(%error, "Failed periodic KV DC Relay membership reconciliation"), }, } } @@ -277,6 +776,49 @@ async fn run_membership_watch( } } +fn publish_membership_if_changed(sender: &watch::Sender, next: DcMembershipView) { + sender.send_if_modified(move |current| { + if current == &next { + return false; + } + *current = next; + true + }); +} + +async fn list_queries( + discovery: &Arc, + queries: &[DiscoveryQuery], +) -> anyhow::Result> { + let results = try_join_all(queries.iter().cloned().map(|query| discovery.list(query))).await?; + Ok(results.into_iter().flatten().collect()) +} + +type RebindingDiscoveryStream = + Pin>> + Send>>; + +async fn open_query_streams( + discovery: &Arc, + queries: &[DiscoveryQuery], + cancel: &CancellationToken, +) -> anyhow::Result> { + let opened = try_join_all( + queries + .iter() + .cloned() + .map(|query| discovery.list_and_watch(query, Some(cancel.clone()))), + ) + .await?; + let mut streams = futures::stream::SelectAll::new(); + for stream in opened { + let stream = stream + .map(|event| event.map(Some)) + .chain(futures::stream::once(async { Ok(None) })); + streams.push(Box::pin(stream) as RebindingDiscoveryStream); + } + Ok(streams) +} + async fn retry_or_cancel(delay: Duration, cancel: &CancellationToken) -> bool { tokio::select! { _ = cancel.cancelled() => false, @@ -284,12 +826,21 @@ async fn retry_or_cancel(delay: Duration, cancel: &CancellationToken) -> bool { } } -fn decode_card(instance: DiscoveryInstance) -> Option<(ModelCardInstanceId, ModelDeploymentCard)> { +fn decode_card(instance: DiscoveryInstance) -> Option<(ModelCardInstanceId, StoredModelCard)> { let DiscoveryInstanceId::Model(id) = instance.id() else { return None; }; + let DiscoveryInstance::Model { card_json, .. } = &instance else { + return None; + }; match instance.deserialize_model::() { - Ok(card) => Some((id, card)), + Ok(card) => Some(( + id, + StoredModelCard { + card, + serialized: card_json.clone(), + }, + )), Err(error) => { tracing::warn!(instance = %id.to_path(), %error, "Ignoring malformed KV DC Relay model card"); None @@ -311,14 +862,39 @@ fn sorted(values: HashSet) -> Vec { values } -fn same_membership(left: &EndpointMembership, right: &EndpointMembership) -> bool { - left.endpoint == right.endpoint - && left.domain == right.domain - && left.compatibility_conflict == right.compatibility_conflict - && left.models == right.models - && left.aliases == right.aliases - && left.roles == right.roles - && left.runtime_configs == right.runtime_configs +fn valid_aliases( + id: &ModelCardInstanceId, + card: &ModelDeploymentCard, + model: &CanonicalModelId, + endpoint: &EndpointId, + invalid_aliases: &mut HashSet<(ModelCardInstanceId, String)>, + warned_invalid_aliases: &HashSet<(ModelCardInstanceId, String)>, +) -> Vec { + let mut aliases = HashSet::new(); + for alias in &card.aliases { + match ModelAlias::new(alias.clone()) { + Ok(alias) if alias.as_str() != model.as_str() => { + aliases.insert(alias); + } + Ok(_) => {} + Err(error) => { + let key = (id.clone(), alias.clone()); + invalid_aliases.insert(key.clone()); + if !warned_invalid_aliases.contains(&key) { + tracing::warn!( + endpoint = %endpoint, + model = %model, + alias, + %error, + "Ignoring invalid model alias" + ); + } + } + } + } + let mut aliases: Vec<_> = aliases.into_iter().collect(); + aliases.sort_unstable(); + aliases } #[cfg(test)] @@ -335,6 +911,159 @@ mod tests { card } + #[test] + fn discovery_scope_requires_explicit_namespaces_or_watch_all() { + let config = KvDcRelayDiscoveryConfig::default(); + assert!(config.validate().is_err()); + + let config = KvDcRelayDiscoveryConfig { + watch_all: true, + ..Default::default() + }; + assert_eq!(config.queries(), vec![DiscoveryQuery::AllModels]); + assert!(config.validate().is_ok()); + } + + #[test] + fn discovery_scope_uses_one_server_side_query_per_namespace() { + let config = KvDcRelayDiscoveryConfig { + namespaces: vec!["prod-a".into(), "prod-b".into()], + endpoint_prefixes: vec!["prod-a.backend".into()], + watch_all: false, + }; + assert!(config.validate().is_ok()); + assert_eq!( + config.queries(), + vec![ + DiscoveryQuery::NamespacedModels { + namespace: "prod-a".into() + }, + DiscoveryQuery::NamespacedModels { + namespace: "prod-b".into() + }, + ] + ); + let filter = config.filter(); + assert!(filter.matches(&EndpointId::from("prod-a.backend.generate"))); + assert!(!filter.matches(&EndpointId::from("prod-a.backend2.generate"))); + assert!(!filter.matches(&EndpointId::from("prod-b.backend.generate"))); + } + + #[test] + fn discovery_scope_rejects_prefix_outside_assigned_namespaces() { + let config = KvDcRelayDiscoveryConfig { + namespaces: vec!["prod-a".into()], + endpoint_prefixes: vec!["prod-b.backend".into()], + watch_all: false, + }; + assert!(config.validate().is_err()); + } + + #[test] + fn discovery_scope_rejects_surrounding_whitespace() { + let padded_namespace = KvDcRelayDiscoveryConfig { + namespaces: vec![" prod-a".into()], + endpoint_prefixes: Vec::new(), + watch_all: false, + }; + assert!(padded_namespace.validate().is_err()); + + let padded_prefix = KvDcRelayDiscoveryConfig { + namespaces: vec!["prod-a".into()], + endpoint_prefixes: vec!["prod-a.backend ".into()], + watch_all: false, + }; + assert!(padded_prefix.validate().is_err()); + } + + #[test] + fn unchanged_membership_does_not_advance_the_watch_version() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + let first = instance("generate", 1, None, card("llama", "meta/llama", 64)); + assert!(state.apply(DiscoveryEvent::Added(first.clone()), &filter)); + let initial = state.view(&filter); + let (sender, mut receiver) = watch::channel(initial.clone()); + let projection_count = state.projection_count; + + let changed = state.apply(DiscoveryEvent::Added(first.clone()), &filter); + if changed { + publish_membership_if_changed(&sender, state.view(&filter)); + } + assert!(!changed); + assert_eq!(state.projection_count, projection_count); + assert!(!receiver.has_changed().unwrap()); + + let changed = state.replace_all(vec![first], &filter); + if changed { + publish_membership_if_changed(&sender, state.view(&filter)); + } + assert!(!changed); + assert_eq!(state.projection_count, projection_count); + assert!(!receiver.has_changed().unwrap()); + + let changed = state.apply( + DiscoveryEvent::Removed(DiscoveryInstanceId::Model(ModelCardInstanceId { + namespace: "prod".to_string(), + component: "backend".to_string(), + endpoint: "generate".to_string(), + instance_id: 999, + model_suffix: None, + })), + &filter, + ); + if changed { + publish_membership_if_changed(&sender, state.view(&filter)); + } + assert!(!changed); + assert_eq!(state.projection_count, projection_count); + assert!(!receiver.has_changed().unwrap()); + + let changed = state.apply( + DiscoveryEvent::Added(instance( + "generate", + 2, + None, + card("llama", "meta/llama", 64), + )), + &filter, + ); + if changed { + publish_membership_if_changed(&sender, state.view(&filter)); + } + assert!(changed); + assert_eq!(state.projection_count, projection_count + 1); + assert!(receiver.has_changed().unwrap()); + assert_ne!(*receiver.borrow_and_update(), initial); + } + + #[test] + fn reappearing_endpoint_advances_its_generation() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + let discovery_instance = instance("generate", 1, None, card("llama", "meta/llama", 64)); + let instance_id = DiscoveryInstanceId::Model(ModelCardInstanceId { + namespace: "prod".to_string(), + component: "backend".to_string(), + endpoint: "generate".to_string(), + instance_id: 1, + model_suffix: None, + }); + + state.apply(DiscoveryEvent::Added(discovery_instance.clone()), &filter); + let initial = state.view(&filter); + assert_eq!(initial.endpoints.len(), 1); + assert_eq!(initial.endpoints.values().next().unwrap().generation, 1); + + state.apply(DiscoveryEvent::Removed(instance_id), &filter); + assert!(state.view(&filter).endpoints.is_empty()); + + state.apply(DiscoveryEvent::Added(discovery_instance), &filter); + let reappeared = state.view(&filter); + assert_eq!(reappeared.endpoints.len(), 1); + assert_eq!(reappeared.endpoints.values().next().unwrap().generation, 2); + } + fn instance( endpoint: &str, instance_id: u64, @@ -351,8 +1080,15 @@ mod tests { } } + fn membership_for_endpoint<'a>( + view: &'a DcMembershipView, + endpoint: &EndpointId, + ) -> &'a EndpointMembership { + &view.endpoints[endpoint] + } + #[test] - fn exact_endpoint_membership_fences_conflicting_base_cards_but_ignores_lora_domains() { + fn incompatible_domains_under_one_endpoint_are_fenced_together() { let filter = DcDiscoveryFilter::default(); let mut state = MembershipState::default(); state.apply( @@ -366,76 +1102,299 @@ mod tests { ); state.apply( DiscoveryEvent::Added(instance( - "embeddings", + "generate", 2, None, card("embed", "nvidia/embed", 32), )), &filter, ); + + let endpoint = EndpointId::from("prod.backend.generate"); + let view = state.view(&filter); + assert_eq!(view.endpoints.len(), 1); + let membership = membership_for_endpoint(&view, &endpoint); + assert!(membership.domain.is_none()); + assert!(!membership.is_materializable()); + assert_eq!(membership.models, ["embed", "llama"]); + assert!(membership.conflicts.iter().any(|conflict| { + matches!( + conflict, + MaterializationConflict::Endpoint { endpoint: conflicted, .. } if conflicted == &endpoint + ) + })); + let conflicted_generation = membership.generation; + + state.apply( + DiscoveryEvent::Removed(DiscoveryInstanceId::Model(ModelCardInstanceId { + namespace: "prod".to_string(), + component: "backend".to_string(), + endpoint: "generate".to_string(), + instance_id: 2, + model_suffix: None, + })), + &filter, + ); + let view = state.view(&filter); + let membership = membership_for_endpoint(&view, &endpoint); + assert_eq!(membership.models, ["llama"]); + assert!(membership.domain.is_some()); + assert!(membership.conflicts.is_empty()); + assert!(membership.is_materializable()); + assert!(membership.generation > conflicted_generation); + } + + #[test] + fn different_base_models_under_one_endpoint_are_a_hard_conflict() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); state.apply( DiscoveryEvent::Added(instance( "generate", - 4, + 1, + None, + card("llama", "meta/llama", 64), + )), + &filter, + ); + state.apply( + DiscoveryEvent::Added(instance( + "generate", + 2, None, - card("llama-public", "meta/llama", 64), + card("chat", "meta/llama", 64), )), &filter, ); + let endpoint = EndpointId::from("prod.backend.generate"); let view = state.view(&filter); - assert_eq!(view.endpoints.len(), 2); - let generate = &view.endpoints[&EndpointId::from("prod.backend.generate")]; + let membership = membership_for_endpoint(&view, &endpoint); + assert!(membership.domain.is_some()); + assert!(!membership.is_materializable()); + assert!(membership.conflicts.iter().any(|conflict| { + matches!( + conflict, + MaterializationConflict::Endpoint { endpoint: conflicted, reason } + if conflicted == &endpoint && reason.contains("canonical base models") + ) + })); + } + + #[test] + fn adapter_is_a_loaded_overlay_on_the_backing_base_domain() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + state.apply( + DiscoveryEvent::Added(instance( + "generate", + 1, + None, + card("llama", "meta/llama", 64), + )), + &filter, + ); + let mut adapter = card("tenant-a", "meta/llama", 64); + adapter.lora = Some(LoraInfo { + name: "tenant-a".to_string(), + max_gpu_lora_count: Some(4), + }); + state.apply( + DiscoveryEvent::Added(instance("generate", 1, Some("tenant-a"), adapter)), + &filter, + ); + + let endpoint = EndpointId::from("prod.backend.generate"); + let view = state.view(&filter); + let membership = membership_for_endpoint(&view, &endpoint); + assert!(membership.is_materializable()); + assert_eq!(membership.runtime_configs.len(), 1); + assert_eq!(membership.models, ["llama", "tenant-a"]); + let adapter = CanonicalModelId::new("tenant-a").unwrap(); + let registration = membership + .registrations + .iter() + .find(|registration| registration.model() == &adapter) + .unwrap(); assert_eq!( - generate.domain.as_ref().unwrap().diagnostic_model_artifact, + registration.target(), + &ModelTarget::Lora { + base_model: CanonicalModelId::new("llama").unwrap(), + adapter, + } + ); + } + + #[test] + fn adapter_metadata_does_not_override_the_base_pool_domain() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + state.apply( + DiscoveryEvent::Added(instance( + "generate", + 1, + None, + card("llama", "meta/llama", 64), + )), + &filter, + ); + let mut adapter = card("tenant-a", "unrelated/adapter", 1); + adapter.lora = Some(LoraInfo { + name: "tenant-a".to_string(), + max_gpu_lora_count: Some(4), + }); + state.apply( + DiscoveryEvent::Added(instance("generate", 1, Some("tenant-a"), adapter)), + &filter, + ); + + let endpoint = EndpointId::from("prod.backend.generate"); + let view = state.view(&filter); + let membership = membership_for_endpoint(&view, &endpoint); + assert!(membership.is_materializable()); + assert!(membership.conflicts.is_empty()); + assert_eq!( + membership + .domain + .as_ref() + .unwrap() + .diagnostic_model_artifact, "meta/llama" ); - assert!(!generate.compatibility_conflict); - assert_eq!(generate.runtime_configs.len(), 2); - assert_eq!(generate.models, vec!["llama", "llama-public"]); + assert_eq!(membership.models, ["llama", "tenant-a"]); + } + #[test] + fn orphan_adapter_is_soft_and_does_not_block_a_valid_base_endpoint() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); state.apply( DiscoveryEvent::Added(instance( "generate", - 3, + 1, None, - card("other", "other/artifact", 64), + card("llama", "meta/llama", 64), )), &filter, ); + let mut adapter = card("tenant-a", "meta/llama", 64); + adapter.lora = Some(LoraInfo { + name: "tenant-a".to_string(), + max_gpu_lora_count: Some(4), + }); + state.apply( + DiscoveryEvent::Added(instance("generate", 2, Some("tenant-a"), adapter)), + &filter, + ); + + let endpoint = EndpointId::from("prod.backend.generate"); let view = state.view(&filter); - let generate = &view.endpoints[&EndpointId::from("prod.backend.generate")]; - assert!(generate.compatibility_conflict); - assert!(generate.domain.is_none()); + let membership = membership_for_endpoint(&view, &endpoint); + assert!(membership.is_materializable()); + assert!(membership.conflicts.is_empty()); + assert_eq!(membership.models, ["llama"]); + } + #[test] + fn base_model_aliases_remain_attached_to_their_canonical_registration() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + let mut model = card("llama", "meta/llama", 64); + model.aliases = vec!["chat".to_string(), "instruct".to_string()]; state.apply( - DiscoveryEvent::Removed(DiscoveryInstanceId::Model(ModelCardInstanceId { - namespace: "prod".to_string(), - component: "backend".to_string(), - endpoint: "generate".to_string(), - instance_id: 3, - model_suffix: None, - })), + DiscoveryEvent::Added(instance("generate", 1, None, model)), &filter, ); - let mut adapter = card("llama-adapter", "unrelated/adapter", 1); + + let view = state.view(&filter); + let endpoint = EndpointId::from("prod.backend.generate"); + let membership = membership_for_endpoint(&view, &endpoint); + assert_eq!(membership.registrations.len(), 1); + let registration = &membership.registrations[0]; + assert_eq!(registration.model().as_str(), "llama"); + assert_eq!( + registration + .aliases() + .iter() + .map(ModelAlias::as_str) + .collect::>(), + vec!["chat", "instruct"] + ); + } + + #[test] + fn request_names_are_scoped_to_each_endpoint_pool() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + let mut fast = card("llama", "meta/llama", 64); + fast.aliases = vec!["chat".to_string()]; + let mut slow = card("mistral", "mistral/model", 64); + slow.aliases = vec!["chat".to_string()]; + state.apply( + DiscoveryEvent::Added(instance("fast", 1, None, fast)), + &filter, + ); + state.apply( + DiscoveryEvent::Added(instance("slow", 2, None, slow)), + &filter, + ); + + let view = state.view(&filter); + assert_eq!(view.endpoints.len(), 2); + assert!(view.endpoints.values().all(|membership| { + membership.is_materializable() && membership.aliases == ["chat"] + })); + } + + #[test] + fn request_names_must_be_unambiguous_within_an_endpoint_pool() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + let mut base = card("llama", "meta/llama", 64); + base.aliases = vec!["tenant-a".to_string()]; + let mut adapter = card("tenant-a", "unrelated/adapter", 1); adapter.lora = Some(LoraInfo { name: "tenant-a".to_string(), - max_gpu_lora_count: None, + max_gpu_lora_count: Some(4), }); + state.apply( + DiscoveryEvent::Added(instance("generate", 1, None, base)), + &filter, + ); state.apply( DiscoveryEvent::Added(instance("generate", 1, Some("tenant-a"), adapter)), &filter, ); let view = state.view(&filter); - let generate = &view.endpoints[&EndpointId::from("prod.backend.generate")]; - assert!(!generate.compatibility_conflict); - assert_eq!( - generate.domain.as_ref().unwrap().diagnostic_model_artifact, - "meta/llama" + let membership = &view.endpoints[&EndpointId::from("prod.backend.generate")]; + assert!(!membership.is_materializable()); + assert!(membership.conflicts.iter().any(|conflict| { + matches!( + conflict, + MaterializationConflict::Endpoint { reason, .. } + if reason.contains("tenant-a") + ) + })); + } + + #[test] + fn invalid_alias_is_soft_and_does_not_block_the_endpoint() { + let filter = DcDiscoveryFilter::default(); + let mut state = MembershipState::default(); + let mut llama = card("llama", "meta/llama", 64); + llama.aliases = vec!["".to_string()]; + state.apply( + DiscoveryEvent::Added(instance("generate", 1, None, llama)), + &filter, ); - assert_eq!(generate.runtime_configs.len(), 2); - assert!(generate.models.iter().any(|model| model == "llama-adapter")); + + let view = state.view(&filter); + let endpoint = EndpointId::from("prod.backend.generate"); + let membership = membership_for_endpoint(&view, &endpoint); + assert!(membership.is_materializable()); + assert_eq!(membership.models, ["llama"]); + assert!(membership.aliases.is_empty()); + assert!(membership.conflicts.is_empty()); } } diff --git a/lib/llm/src/kv_dc_relay/host.rs b/lib/llm/src/kv_dc_relay/host.rs index 8aff987bf396..36368c50daab 100644 --- a/lib/llm/src/kv_dc_relay/host.rs +++ b/lib/llm/src/kv_dc_relay/host.rs @@ -1,13 +1,11 @@ // SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. // SPDX-License-Identifier: Apache-2.0 -//! DC-scoped KV-cache Relay with one serialized CKF actor per serving endpoint. +//! DC-scoped KV-cache Relay with one serialized CKF actor per physical pool. //! //! Dynamo discovery and worker-local recovery feed endpoint actors. The actors' //! exact member ownership is authoritative; each materialization publishes one -//! physical CKF layout for a future Relay-to-global-router adapter. -//! The subscription seam remains crate-private: a standalone/WAN publisher API -//! requires delivery cursors and recovery semantics and is intentionally deferred. +//! physical CKF layout through the pool catalog subscription boundary. //! //! NOTE: One serialized actor per endpoint pool is the current measured choice, not a claim that //! it scales indefinitely. A worker-partitioned, multi-issuer Mooncake comparison found the @@ -15,34 +13,38 @@ //! dedicated Relay campaign before changing this ownership model; further producer optimization //! will likely be needed for substantially larger DC-scale pools. -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; +use std::future::Future; use std::sync::Arc; use std::time::Duration; -#[cfg(feature = "ckf-diagnostics")] -use std::sync::atomic::AtomicU64; #[cfg(feature = "ckf-diagnostics")] use std::sync::atomic::Ordering; -#[cfg(feature = "ckf-diagnostics")] -use std::time::Instant; use dynamo_kv_router::identity::PoolId; -use dynamo_kv_router::indexer::cuckoo::{CkfConfig, CkfFailureAction}; +use dynamo_kv_router::indexer::cuckoo::CkfFailureAction; use dynamo_kv_router::protocols::{DpRank, KvCacheEventError, WorkerId}; use dynamo_runtime::component::Component; use dynamo_runtime::protocols::EndpointId; use dynamo_runtime::traits::DistributedRuntimeProvider; use parking_lot::Mutex; +use rand::TryRngCore; use serde::Serialize; -use tokio::sync::{RwLock, Semaphore, mpsc, watch}; -use tokio::task::JoinHandle; +use tokio::sync::{RwLock, Semaphore, watch}; +use tokio::task::{JoinHandle, JoinSet}; use tokio_util::sync::CancellationToken; -use super::actor::{ActorFault, KvDcRelayHandle, KvDcRelayRecoveryTarget, StreamScope}; +use super::actor::{ActorFault, DEFAULT_FAULT_CAPACITY, KvDcRelayHandle, KvDcRelayRecoveryTarget}; use super::discovery::{ - DcDiscoveryFilter, DcMembershipView, DcMembershipWatch, EndpointMembership, KvCacheDomainKey, + DcMembershipView, DcMembershipWatch, EndpointMembership, KvCacheDomainKey, + KvDcRelayDiscoveryConfig, MaterializationConflict, +}; +use super::identity::{CanonicalModelRegistration, DcPoolCatalog, DcRelayIdentity}; +use super::pool_registry::{ + PoolActorConfig, PoolAttachRequest, PoolAttachment, PoolRegistry, PoolRetirementMode, + drain_faults_while, }; -use super::resolution::{EndpointLocator, PoolBinding, stable_dc_id}; +use super::resolution::stable_dc_id; use crate::discovery::{KvSourceMembershipCoordinator, KvSourceMembershipWatch}; #[cfg(feature = "ckf-diagnostics")] use crate::kv_router::indexer::WorkerQueryHealthSnapshot; @@ -50,6 +52,7 @@ use crate::kv_router::indexer::{ DEFAULT_RECOVERY_ATTEMPT_TIMEOUT, RecoverySupervisor, TargetFaultDisposition, start_target_subscriber, }; +use crate::local_model::runtime_config::ModelRuntimeConfig; pub const DEFAULT_EXPECTED_UNIQUE_BLOCKS: usize = 1_048_576; const DEFAULT_RECOVERY_FETCH_CONCURRENCY: usize = 16; @@ -129,7 +132,8 @@ pub struct KvDcRelayStats { #[non_exhaustive] pub struct KvDcRelayIdentityStats { pub dc_id: String, - pub process_incarnation: u64, + pub drt_instance_id: u64, + pub relay_incarnation: u64, } #[cfg(feature = "ckf-diagnostics")] @@ -140,7 +144,7 @@ pub struct KvDcRelayEndpointStats { pub lifecycle: String, pub layout_generation: u64, pub cache_domain: Option, - pub compatibility_conflict: bool, + pub membership_conflicts: Vec, pub models: Vec, pub aliases: Vec, pub roles: Vec, @@ -252,7 +256,8 @@ pub struct KvDcRelayHealth { #[derive(Debug, Clone, Serialize)] #[non_exhaustive] pub struct KvDcRelayDiagnosticSnapshot { - pub process_incarnation: u64, + pub drt_instance_id: u64, + pub relay_incarnation: u64, pub dc_id: String, pub serving_endpoint: String, pub layout_generation: u64, @@ -322,11 +327,11 @@ struct EndpointSlotTask { task: JoinHandle<()>, } -struct EndpointActorRuntime { - handle: KvDcRelayHandle, +struct EndpointPoolRuntime { + attachment: PoolAttachment, recovery: RecoverySupervisor, - faults: mpsc::Receiver, binding: ActorBinding, + registrations: Vec, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -335,16 +340,135 @@ struct ActorBinding { kv_state_endpoint: EndpointId, } +const MAX_PENDING_SOURCE_FAULTS: usize = DEFAULT_FAULT_CAPACITY; + +enum ProducerFenceTrigger { + Fault(ActorFault), + PendingOverflow(ActorFault), +} + +enum PendingActorAction { + Fault(ActorFault), + ProducerFence(ProducerFenceTrigger), +} + +#[derive(Default)] +struct PendingActorFaults { + producer_fence: Option, + source_faults: HashMap<(WorkerId, DpRank), ActorFault>, + source_order: VecDeque<(WorkerId, DpRank)>, +} + +impl PendingActorFaults { + fn push(&mut self, fault: ActorFault) { + if self.producer_fence.is_some() { + return; + } + + match fault.disposition.action { + CkfFailureAction::ContinueCapacityOmission => {} + CkfFailureAction::ReportResourceFailure | CkfFailureAction::RejectSource => { + self.push_source_fault(fault); + } + CkfFailureAction::FenceAndRebuildProducer => { + self.install_producer_fence(ProducerFenceTrigger::Fault(fault)); + } + CkfFailureAction::DeactivateAndSnapshot | CkfFailureAction::RetrySnapshot => { + unreachable!("consumer-lane disposition cannot originate from Relay actor") + } + } + } + + fn push_source_fault(&mut self, fault: ActorFault) { + let key = (fault.worker_id, fault.dp_rank); + if let Some(current) = self.source_faults.get_mut(&key) { + let current_epoch = current.source_epoch.get(); + let candidate_epoch = fault.source_epoch.get(); + if candidate_epoch > current_epoch + || (candidate_epoch == current_epoch + && is_stronger_source_fault( + fault.disposition.action, + current.disposition.action, + )) + { + *current = fault; + } + return; + } + + if self.source_faults.len() >= MAX_PENDING_SOURCE_FAULTS { + self.install_producer_fence(ProducerFenceTrigger::PendingOverflow(fault)); + return; + } + + self.source_order.push_back(key); + self.source_faults.insert(key, fault); + } + + fn install_producer_fence(&mut self, trigger: ProducerFenceTrigger) { + self.source_faults.clear(); + self.source_order.clear(); + self.producer_fence = Some(trigger); + } + + fn drain_ready(&mut self, receiver: &mut tokio::sync::mpsc::Receiver) { + while self.producer_fence.is_none() { + let Ok(fault) = receiver.try_recv() else { + break; + }; + self.push(fault); + } + } + + fn pop_front(&mut self) -> Option { + if let Some(trigger) = self.producer_fence.take() { + return Some(PendingActorAction::ProducerFence(trigger)); + } + while let Some(key) = self.source_order.pop_front() { + if let Some(fault) = self.source_faults.remove(&key) { + return Some(PendingActorAction::Fault(fault)); + } + } + None + } + + fn take_producer_fence(&mut self) -> Option { + self.producer_fence.take() + } + + fn clear(&mut self) { + self.producer_fence = None; + self.source_faults.clear(); + self.source_order.clear(); + } + + #[cfg(test)] + fn len(&self) -> usize { + usize::from(self.producer_fence.is_some()) + self.source_faults.len() + } +} + +fn is_stronger_source_fault(candidate: CkfFailureAction, current: CkfFailureAction) -> bool { + matches!( + (current, candidate), + ( + CkfFailureAction::ReportResourceFailure, + CkfFailureAction::RejectSource + ) + ) +} + /// DC-wide Relay host. It is intentionally not scoped to a model, namespace, or endpoint. pub struct KvDcRelay { #[cfg(feature = "ckf-diagnostics")] dc_id: Arc, #[cfg(feature = "ckf-diagnostics")] - process_incarnation: u64, + relay_identity: DcRelayIdentity, cancel: CancellationToken, membership: Mutex>, supervisor: Mutex>>, statuses: Arc>>, + pools: Arc, } impl KvDcRelay { @@ -353,6 +477,11 @@ impl KvDcRelay { dc_id: String, config: KvDcRelayConfig, ) -> anyhow::Result { + anyhow::ensure!(!dc_id.is_empty(), "KV DC Relay dc_id must not be empty"); + anyhow::ensure!( + dc_id.trim() == dc_id, + "KV DC Relay dc_id must not contain leading or trailing whitespace" + ); anyhow::ensure!( config.publication_threshold != 0, "KV DC Relay publication_threshold must be positive" @@ -369,27 +498,35 @@ impl KvDcRelay { threshold: config.publication_threshold, delay: Duration::from_millis(config.publication_delay_ms), }; + let discovery = KvDcRelayDiscoveryConfig { + watch_all: config.namespace_filter.is_none(), + namespaces: config.namespace_filter.into_iter().collect(), + endpoint_prefixes: config.endpoint_prefix.into_iter().collect(), + }; let cancel = component.drt().child_token(); - let membership = DcMembershipWatch::start( - component.drt().discovery(), - DcDiscoveryFilter { - namespace: config.namespace_filter, - endpoint_prefix: config.endpoint_prefix, - }, - cancel.clone(), - ) - .await?; + let membership = + DcMembershipWatch::start(component.drt().discovery(), discovery, cancel.clone()) + .await?; let membership_rx = membership.subscribe(); let statuses = Arc::new(RwLock::new(HashMap::new())); let dc_id: Arc = Arc::from(dc_id); - let process_incarnation = component.drt().connection_id(); + let relay_identity = + DcRelayIdentity::new(component.drt().connection_id(), new_relay_incarnation()?); + let ckf_dc_id = stable_dc_id(dc_id.as_ref()); + let pools = Arc::new(PoolRegistry::new( + relay_identity, + PoolActorConfig { + expected_unique_blocks: DEFAULT_EXPECTED_UNIQUE_BLOCKS, + publication_threshold: publication.threshold, + publication_delay: publication.delay, + }, + )); let supervisor = tokio::spawn(run_host_supervisor( component, - dc_id.clone(), - process_incarnation, + ckf_dc_id, membership_rx, statuses.clone(), - publication, + pools.clone(), Duration::from_millis(config.recovery_attempt_timeout_ms), cancel.child_token(), )); @@ -397,14 +534,23 @@ impl KvDcRelay { #[cfg(feature = "ckf-diagnostics")] dc_id, #[cfg(feature = "ckf-diagnostics")] - process_incarnation, + relay_identity, cancel, membership: Mutex::new(Some(membership)), supervisor: Mutex::new(Some(supervisor)), statuses, + pools, }) } + pub fn pool_catalog(&self) -> DcPoolCatalog { + self.pools.catalog() + } + + pub fn watch_pool_catalog(&self) -> watch::Receiver { + self.pools.watch_catalog() + } + #[cfg(feature = "ckf-diagnostics")] pub async fn stats(&self) -> Result { let statuses: Vec<_> = self @@ -412,18 +558,19 @@ impl KvDcRelay { .read() .await .iter() - .map(|(endpoint, status)| (endpoint.clone(), status.clone())) + .map(|(slot_id, status)| (slot_id.clone(), status.clone())) .collect(); let mut endpoints = Vec::with_capacity(statuses.len()); - for (endpoint, status) in statuses { - endpoints.push(endpoint_stats(endpoint, status).await?); + for (slot_id, status) in statuses { + endpoints.push(endpoint_stats(slot_id, status).await?); } endpoints .sort_unstable_by(|left, right| left.serving_endpoint.cmp(&right.serving_endpoint)); Ok(KvDcRelayStats { identity: KvDcRelayIdentityStats { dc_id: self.dc_id.to_string(), - process_incarnation: self.process_incarnation, + drt_instance_id: self.relay_identity.drt_instance_id(), + relay_incarnation: self.relay_identity.relay_incarnation(), }, endpoints, }) @@ -452,7 +599,8 @@ impl KvDcRelay { let format = actor_snapshot.identity.format(); let aggregation = actor_snapshot.stats.aggregation(); Ok(KvDcRelayDiagnosticSnapshot { - process_incarnation: self.process_incarnation, + drt_instance_id: self.relay_identity.drt_instance_id(), + relay_incarnation: self.relay_identity.relay_incarnation(), dc_id: self.dc_id.to_string(), serving_endpoint: endpoint.to_string(), layout_generation, @@ -518,38 +666,43 @@ impl KvDcRelay { } } +fn new_relay_incarnation() -> anyhow::Result { + let random_id = rand::rngs::OsRng + .try_next_u64() + .map_err(|error| anyhow::anyhow!("failed to generate Relay incarnation: {error}"))?; + Ok(random_id & (i64::MAX as u64)) +} + #[allow(clippy::too_many_arguments)] async fn run_host_supervisor( component: Component, - dc_id: Arc, - process_incarnation: u64, + ckf_dc_id: dynamo_kv_router::DcId, mut membership_rx: watch::Receiver, statuses: Arc>>, - publication: ActorPublicationConfig, + pools: Arc, recovery_attempt_timeout: Duration, cancel: CancellationToken, ) { let recovery_fetch_permit = Arc::new(Semaphore::new(DEFAULT_RECOVERY_FETCH_CONCURRENCY)); let mut slots: HashMap = HashMap::new(); - let ckf_dc_id = stable_dc_id(dc_id.as_ref()); + let mut retired_slots = JoinSet::new(); loop { let mut view = membership_rx.borrow_and_update().clone(); reject_duplicate_live_pools(&mut view, ckf_dc_id); - for (endpoint, membership) in &view.endpoints { - let slot = slots.entry(endpoint.clone()).or_insert_with(|| { + for (slot_id, membership) in view.endpoints.iter() { + let slot = slots.entry(slot_id.clone()).or_insert_with(|| { let (metadata, metadata_rx) = watch::channel(None); let status = Arc::new(RwLock::new(EndpointSlotStatus::default())); let task = tokio::spawn(run_endpoint_slot( component.clone(), - dc_id.clone(), - process_incarnation, - endpoint.clone(), + ckf_dc_id, + slot_id.clone(), metadata_rx, status.clone(), Arc::new(Semaphore::new(1)), recovery_fetch_permit.clone(), - publication, + pools.clone(), recovery_attempt_timeout, cancel.child_token(), )); @@ -559,16 +712,12 @@ async fn run_host_supervisor( task, } }); - slot.metadata.send_replace(Some(membership.clone())); - } - for (endpoint, slot) in &slots { - if !view.endpoints.contains_key(endpoint) { - slot.metadata.send_replace(None); - } + publish_endpoint_metadata_if_changed(&slot.metadata, membership); } + retire_departed_endpoint_slots(&view, &mut slots, &mut retired_slots); *statuses.write().await = slots .iter() - .map(|(endpoint, slot)| (endpoint.clone(), slot.status.clone())) + .map(|(slot_id, slot)| (slot_id.clone(), slot.status.clone())) .collect(); tokio::select! { @@ -578,29 +727,26 @@ async fn run_host_supervisor( break; } } + retired = retired_slots.join_next(), if !retired_slots.is_empty() => { + report_retired_endpoint_slot(retired); + } } } - drop( - slots - .values() - .map(|slot| slot.metadata.clone()) - .collect::>(), - ); - for (_, slot) in slots { + for (slot_id, slot) in slots { drop(slot.metadata); - if let Err(error) = slot.task.await - && !error.is_cancelled() - { - tracing::warn!(%error, "KV DC Relay endpoint slot failed during shutdown"); - } + report_endpoint_slot_exit(slot_id, slot.task.await); + } + while let Some(retired) = retired_slots.join_next().await { + report_retired_endpoint_slot(Some(retired)); } + pools.shutdown().await; } fn reject_duplicate_live_pools(view: &mut DcMembershipView, dc_id: dynamo_kv_router::DcId) { let mut owners: HashMap> = HashMap::new(); - for (endpoint, membership) in &view.endpoints { - if membership.compatibility_conflict { + for (endpoint, membership) in view.endpoints.iter() { + if !membership.conflicts.is_empty() { continue; } let Some(domain) = &membership.domain else { @@ -612,43 +758,148 @@ fn reject_duplicate_live_pools(view: &mut DcMembershipView, dc_id: dynamo_kv_rou .push(endpoint.clone()); } - for (pool_id, endpoints) in owners { + for (pool_id, mut endpoints) in owners { if endpoints.len() < 2 { continue; } + endpoints.sort_unstable_by_key(ToString::to_string); tracing::error!( %pool_id, endpoints = ?endpoints, "multiple live serving endpoints resolve to one CKF pool; fencing all colliding endpoints" ); - for endpoint in endpoints { - if let Some(membership) = view.endpoints.get_mut(&endpoint) { - membership.compatibility_conflict = true; - } + let memberships = Arc::make_mut(&mut view.endpoints); + for endpoint in &endpoints { + let Some(membership) = memberships.get_mut(endpoint) else { + continue; + }; + membership + .conflicts + .push(MaterializationConflict::Endpoint { + endpoint: endpoint.clone(), + reason: format!("pool {pool_id} is claimed by multiple serving endpoints"), + }); + } + } +} + +fn inactive_slot_lifecycle(membership: Option<&EndpointMembership>) -> SlotLifecycle { + match membership { + None => SlotLifecycle::Lightweight, + Some(membership) if !membership.conflicts.is_empty() => SlotLifecycle::Fenced, + Some(_) => SlotLifecycle::Discovered, + } +} + +fn publish_endpoint_metadata_if_changed( + sender: &watch::Sender>, + membership: &EndpointMembership, +) { + sender.send_if_modified(|current| { + if current.as_ref() == Some(membership) { + return false; } + *current = Some(membership.clone()); + true + }); +} + +fn retire_departed_endpoint_slots( + view: &DcMembershipView, + slots: &mut HashMap, + retired_slots: &mut JoinSet<(EndpointId, Result<(), tokio::task::JoinError>)>, +) { + let departed: Vec<_> = slots + .keys() + .filter(|slot_id| !view.endpoints.contains_key(*slot_id)) + .cloned() + .collect(); + for slot_id in departed { + let Some(slot) = slots.remove(&slot_id) else { + continue; + }; + drop(slot.metadata); + retired_slots.spawn(async move { + let result = slot.task.await; + (slot_id, result) + }); + } +} + +type RetiredEndpointSlot = + Result<(EndpointId, Result<(), tokio::task::JoinError>), tokio::task::JoinError>; + +fn report_retired_endpoint_slot(retired: Option) { + match retired { + Some(Ok((slot_id, result))) => report_endpoint_slot_exit(slot_id, result), + Some(Err(error)) if !error.is_cancelled() => { + tracing::warn!(%error, "KV DC Relay endpoint retirement monitor failed"); + } + Some(Err(_)) | None => {} + } +} + +fn report_endpoint_slot_exit(slot_id: EndpointId, result: Result<(), tokio::task::JoinError>) { + if let Err(error) = result + && !error.is_cancelled() + { + tracing::warn!(endpoint = %slot_id, %error, "KV DC Relay endpoint slot failed"); + } +} + +fn report_actor_fault(endpoint: &EndpointId, fault: &ActorFault) { + tracing::error!( + %endpoint, + worker_id = fault.worker_id, + dp_rank = fault.dp_rank, + event_id = ?fault.event_id, + category = ?fault.category, + error = %fault.message, + "KV DC Relay actor failed an admitted mutation" + ); +} + +fn report_producer_fence_trigger(endpoint: &EndpointId, trigger: &ProducerFenceTrigger) { + match trigger { + ProducerFenceTrigger::Fault(fault) => report_actor_fault(endpoint, fault), + ProducerFenceTrigger::PendingOverflow(fault) => tracing::error!( + %endpoint, + worker_id = fault.worker_id, + dp_rank = fault.dp_rank, + source_epoch = fault.source_epoch.get(), + event_id = ?fault.event_id, + category = ?fault.category, + action = ?fault.disposition.action, + error = %fault.message, + pending_capacity = MAX_PENDING_SOURCE_FAULTS, + "KV DC Relay pending source faults exceeded their bound; fencing producer" + ), } } #[allow(clippy::too_many_arguments)] async fn run_endpoint_slot( component: Component, - dc_id: Arc, - process_incarnation: u64, - endpoint: EndpointId, + ckf_dc_id: dynamo_kv_router::DcId, + slot_id: EndpointId, mut metadata_rx: watch::Receiver>, status: SharedEndpointStatus, rebuild_permit: Arc, recovery_fetch_permit: Arc, - publication: ActorPublicationConfig, + pools: Arc, recovery_attempt_timeout: Duration, cancel: CancellationToken, ) { - let mut config_tx: Option< - watch::Sender>, - > = None; + let endpoint = slot_id.clone(); + let mut config_tx: Option>> = None; let mut source_watch: Option = None; - let mut actor: Option = None; + let mut runtime: Option = None; let mut layout_generation = 0u64; + let mut retry_binding: Option = None; + let mut retry_delay = Duration::from_millis(100); + let mut start_failures = 0u64; + let mut registration_refresh_failures = 0u64; + let mut pending_faults = PendingActorFaults::default(); loop { let membership = metadata_rx.borrow_and_update().clone(); @@ -656,14 +907,20 @@ async fn run_endpoint_slot( let mut current = status.write().await; current.membership = membership.clone(); current.layout_generation = layout_generation; - if membership.is_some() && current.lifecycle == SlotLifecycle::Lightweight { - current.lifecycle = SlotLifecycle::Discovered; + if runtime.is_none() { + current.lifecycle = inactive_slot_lifecycle(membership.as_ref()); } } if let Some(membership) = &membership { if let Some(sender) = &config_tx { - sender.send_replace(membership.runtime_configs.clone()); + sender.send_if_modified(|current| { + if current == &membership.runtime_configs { + return false; + } + current.clone_from(&membership.runtime_configs); + true + }); } else { let (sender, configs) = watch::channel(membership.runtime_configs.clone()); let coordinator = KvSourceMembershipCoordinator::start( @@ -677,8 +934,13 @@ async fn run_endpoint_slot( } let source_view = source_watch.as_ref().map(|watch| watch.borrow().clone()); + let source_binding_pending = membership.as_ref().zip(source_view.as_ref()).is_some_and( + |(membership, source_view)| { + !source_view.matches_binding_inputs(&membership.runtime_configs) + }, + ); let desired_binding = membership.as_ref().and_then(|membership| { - if membership.compatibility_conflict { + if !membership.is_materializable() { return None; } let domain = membership.domain.clone()?; @@ -693,19 +955,33 @@ async fn run_endpoint_slot( kv_state_endpoint, }) }); + if retry_binding.as_ref() != desired_binding.as_ref() { + retry_binding = desired_binding.clone(); + retry_delay = Duration::from_millis(100); + start_failures = 0; + registration_refresh_failures = 0; + } - let binding_changed = actor + let binding_changed = runtime .as_ref() .is_some_and(|active| Some(&active.binding) != desired_binding.as_ref()); if binding_changed || membership.is_none() { - if let Some(active) = actor.take() { - status.write().await.lifecycle = SlotLifecycle::Draining; - stop_endpoint_actor(active).await; + if let Some(active) = runtime.take() { + let lifecycle = inactive_slot_lifecycle(membership.as_ref()); + status.write().await.lifecycle = if lifecycle == SlotLifecycle::Fenced { + SlotLifecycle::Fenced + } else { + SlotLifecycle::Draining + }; + if lifecycle == SlotLifecycle::Fenced { + fence_endpoint_pool(active, &pools).await; + } else { + stop_endpoint_pool(active, &pools).await; + } + pending_faults.clear(); let mut current = status.write().await; current.actor = None; - if membership.is_some() { - current.lifecycle = SlotLifecycle::Discovered; - } + current.lifecycle = lifecycle; } if membership.is_none() { config_tx = None; @@ -719,7 +995,47 @@ async fn run_endpoint_slot( } } - if actor.is_none() + if let (Some(active), Some(membership)) = (runtime.as_mut(), membership.as_ref()) + && membership.is_materializable() + && active.registrations != membership.registrations + { + match refresh_pool_registrations( + &pools, + &mut active.attachment, + &active.binding, + desired_binding.as_ref(), + source_binding_pending, + &membership.registrations, + ) + .await + { + Ok(RegistrationRefresh::Skipped) => {} + Ok(RegistrationRefresh::Published) => { + registration_refresh_failures = 0; + active.registrations.clone_from(&membership.registrations); + } + Err(error) => { + registration_refresh_failures = registration_refresh_failures.saturating_add(1); + if registration_refresh_failures == 1 { + tracing::warn!( + %endpoint, + %error, + "Failed to refresh KV DC Relay model bindings" + ); + } else { + tracing::debug!( + %endpoint, + %error, + registration_refresh_failures, + "KV DC Relay model binding refresh failed again" + ); + } + } + } + } + + if runtime.is_none() + && !source_binding_pending && let (Some(binding), Some(membership), Some(membership_watch)) = ( desired_binding.clone(), membership.clone(), @@ -727,44 +1043,47 @@ async fn run_endpoint_slot( ) { status.write().await.lifecycle = SlotLifecycle::Starting; - let candidate_generation = membership.generation; - layout_generation = layout_generation.saturating_add(1); - match start_endpoint_actor( + match start_endpoint_pool( component.clone(), - dc_id.clone(), - process_incarnation, + ckf_dc_id, endpoint.clone(), - layout_generation, binding.clone(), + membership.registrations.clone(), membership_watch, rebuild_permit.clone(), recovery_fetch_permit.clone(), - publication, + pools.clone(), recovery_attempt_timeout, cancel.child_token(), ) .await { Ok(candidate) - if metadata_rx - .borrow() - .as_ref() - .is_some_and(|current| current.generation == candidate_generation) + if metadata_rx.borrow().as_ref() == Some(&membership) && source_watch.as_ref().and_then(|watch| { watch.borrow().resolved_kv_state_endpoint().cloned() }) == Some(binding.kv_state_endpoint.clone()) => { + retry_delay = Duration::from_millis(100); + start_failures = 0; + registration_refresh_failures = 0; + layout_generation = candidate.attachment.layout_generation; let mut current = status.write().await; current.layout_generation = layout_generation; - current.actor = Some(candidate.handle.clone()); + current.actor = Some(candidate.attachment.handle.clone()); current.lifecycle = SlotLifecycle::Active; - actor = Some(candidate); + runtime = Some(candidate); } Ok(candidate) => { - stop_endpoint_actor(candidate).await; + stop_endpoint_pool(candidate, &pools).await; } Err(error) => { - tracing::error!(%endpoint, %error, "Failed to materialize KV DC Relay endpoint actor"); + start_failures = start_failures.saturating_add(1); + if start_failures == 1 { + tracing::error!(%endpoint, %error, "Failed to materialize KV DC Relay endpoint actor"); + } else { + tracing::debug!(%endpoint, %error, start_failures, retry_ms = retry_delay.as_millis(), "KV DC Relay endpoint actor retry failed"); + } let mut current = status.write().await; current.lifecycle = SlotLifecycle::Fenced; current.actor = None; @@ -772,121 +1091,173 @@ async fn run_endpoint_slot( } } - if membership - .as_ref() - .is_some_and(|membership| membership.compatibility_conflict) - { - status.write().await.lifecycle = SlotLifecycle::Fenced; - } - enum SlotInput { Metadata, Source, - Fault(ActorFault), - ActorExited, + SourceClosed, + Fault(PendingActorAction), + PoolUnavailable, Health, + Retry, Cancelled, } + if let Some(active) = runtime.as_mut() { + pending_faults.drain_ready(&mut active.attachment.faults); + } + let pool_cancel = runtime + .as_ref() + .map(|active| active.attachment.pool_cancel.clone()); let input = tokio::select! { _ = cancel.cancelled() => SlotInput::Cancelled, changed = metadata_rx.changed() => { if changed.is_ok() { SlotInput::Metadata } else { SlotInput::Cancelled } } - changed = async { source_watch.as_mut().expect("guarded source watch").changed().await }, if source_watch.is_some() => { - if changed.is_ok() { SlotInput::Source } else { SlotInput::Metadata } + changed = async { + let Some(source_watch) = source_watch.as_mut() else { + return std::future::pending().await; + }; + source_watch.changed().await + } => { + if changed.is_ok() { SlotInput::Source } else { SlotInput::SourceClosed } } - fault = async { actor.as_mut().expect("guarded actor").faults.recv().await }, if actor.is_some() => { + fault = async { + if let Some(fault) = pending_faults.pop_front() { + return Some(fault); + } + let Some(runtime) = runtime.as_mut() else { + return std::future::pending().await; + }; + runtime + .attachment + .faults + .recv() + .await + .map(PendingActorAction::Fault) + } => { match fault { Some(fault) => SlotInput::Fault(fault), - None => SlotInput::ActorExited, + None => SlotInput::PoolUnavailable, } } - _ = diagnostic_tick(), if actor.is_some() => SlotInput::Health, + _ = async { + let Some(pool_cancel) = pool_cancel.as_ref() else { + return std::future::pending().await; + }; + pool_cancel.cancelled().await + } => SlotInput::PoolUnavailable, + _ = diagnostic_tick(), if runtime.is_some() => SlotInput::Health, + _ = tokio::time::sleep(retry_delay), if runtime.is_none() && desired_binding.is_some() => SlotInput::Retry, }; match input { SlotInput::Metadata | SlotInput::Source | SlotInput::Health => {} - SlotInput::ActorExited => { - tracing::error!(%endpoint, "KV DC Relay actor exited unexpectedly; rebuilding its producer generation"); + SlotInput::SourceClosed => { + tracing::debug!(%endpoint, "KV source membership watch closed; rebinding"); + config_tx = None; + source_watch = None; + } + SlotInput::Retry => { + retry_delay = retry_delay.saturating_mul(2).min(Duration::from_secs(5)); + } + SlotInput::PoolUnavailable => { + tracing::warn!(%endpoint, "KV DC Relay pool generation became unavailable; restarting pool actor"); status.write().await.lifecycle = SlotLifecycle::Fenced; - if let Some(active) = actor.take() { - stop_endpoint_actor(active).await; + if let Some(active) = runtime.take() { + fence_endpoint_pool(active, &pools).await; } + pending_faults.clear(); let mut current = status.write().await; current.actor = None; } - SlotInput::Fault(fault) => { - tracing::error!( - %endpoint, - worker_id = fault.worker_id, - dp_rank = fault.dp_rank, - event_id = ?fault.event_id, - category = ?fault.category, - error = %fault.message, - "KV DC Relay actor failed an admitted mutation" - ); - match fault.disposition.action { - CkfFailureAction::ContinueCapacityOmission => {} - CkfFailureAction::ReportResourceFailure => { - if let Some(active) = actor.as_ref() { - let disposition = active - .recovery - .client() - .handle_target_fault( - fault.worker_id, - fault.dp_rank, - fault.source_epoch, - false, - ) - .await; - if disposition == TargetFaultDisposition::Fenced { - active - .recovery - .client() - .reject_source( - fault.worker_id, - fault.dp_rank, - fault.source_epoch, + SlotInput::Fault(action) => { + let mut retirement_mode = None; + match action { + PendingActorAction::ProducerFence(trigger) => { + report_producer_fence_trigger(&endpoint, &trigger); + retirement_mode = Some(PoolRetirementMode::Fenced); + } + PendingActorAction::Fault(fault) => { + report_actor_fault(&endpoint, &fault); + match fault.disposition.action { + CkfFailureAction::ContinueCapacityOmission => {} + CkfFailureAction::ReportResourceFailure => { + if let Some(active) = runtime.as_mut() { + let client = active.recovery.client().clone(); + match collect_pending_while( + &mut active.attachment.faults, + &mut pending_faults, + client.handle_target_fault( + fault.worker_id, + fault.dp_rank, + fault.source_epoch, + false, + ), ) - .await; - status.write().await.lifecycle = SlotLifecycle::Fenced; + .await + { + FaultCollection::Completed(disposition) => { + retirement_mode = + target_fault_retirement_mode(disposition); + } + FaultCollection::ProducerFence(trigger) => { + report_producer_fence_trigger(&endpoint, &trigger); + retirement_mode = Some(PoolRetirementMode::Fenced); + } + } + } + } + CkfFailureAction::RejectSource => { + if let Some(active) = runtime.as_mut() { + let client = active.recovery.client().clone(); + if let FaultCollection::ProducerFence(trigger) = + collect_pending_while( + &mut active.attachment.faults, + &mut pending_faults, + client.reject_source( + fault.worker_id, + fault.dp_rank, + fault.source_epoch, + ), + ) + .await + { + report_producer_fence_trigger(&endpoint, &trigger); + retirement_mode = Some(PoolRetirementMode::Fenced); + } + } + } + CkfFailureAction::FenceAndRebuildProducer => { + retirement_mode = Some(PoolRetirementMode::Fenced); + } + CkfFailureAction::DeactivateAndSnapshot + | CkfFailureAction::RetrySnapshot => { + unreachable!( + "consumer-lane disposition cannot originate from Relay actor" + ) } } } - CkfFailureAction::RejectSource => { - if let Some(active) = actor.as_ref() { - active - .recovery - .client() - .reject_source(fault.worker_id, fault.dp_rank, fault.source_epoch) - .await; - } - } - CkfFailureAction::FenceAndRebuildProducer => { - // The producer's exact state is suspect. Retire its publisher and source - // bindings before the slot loop constructs a fresh layout generation. - status.write().await.lifecycle = SlotLifecycle::Fenced; - if let Some(active) = actor.take() { - fence_endpoint_actor(active).await; - } - status.write().await.actor = None; - } - CkfFailureAction::DeactivateAndSnapshot | CkfFailureAction::RetrySnapshot => { - unreachable!("consumer-lane disposition cannot originate from Relay actor") + } + if let Some(mode) = retirement_mode { + status.write().await.lifecycle = SlotLifecycle::Fenced; + if let Some(active) = runtime.take() { + retire_endpoint_pool(active, &pools, mode).await; } + pending_faults.clear(); + status.write().await.actor = None; } } SlotInput::Cancelled => break, } #[cfg(feature = "ckf-diagnostics")] - if let Some(active) = &actor { + if let Some(active) = &runtime { status.write().await.recovery = active.recovery.client().health_snapshot().await; } } - if let Some(active) = actor { + if let Some(active) = runtime { status.write().await.lifecycle = SlotLifecycle::Draining; - stop_endpoint_actor(active).await; + stop_endpoint_pool(active, &pools).await; } let mut current = status.write().await; current.actor = None; @@ -901,37 +1272,26 @@ async fn diagnostic_tick() { } #[allow(clippy::too_many_arguments)] -async fn start_endpoint_actor( +async fn start_endpoint_pool( component: Component, - dc_id: Arc, - process_incarnation: u64, + ckf_dc_id: dynamo_kv_router::DcId, endpoint: EndpointId, - layout_generation: u64, binding: ActorBinding, + registrations: Vec, membership_watch: KvSourceMembershipWatch, rebuild_permit: Arc, recovery_fetch_permit: Arc, - publication: ActorPublicationConfig, + pools: Arc, recovery_attempt_timeout: Duration, cancel: CancellationToken, -) -> anyhow::Result { - let mut config = CkfConfig::new(DEFAULT_EXPECTED_UNIQUE_BLOCKS); - config.publish_every_n_events = publication.threshold; - let ckf_dc_id = stable_dc_id(dc_id.as_ref()); - let scope = StreamScope { - process_incarnation, - layout_generation, - pool_binding: PoolBinding::new( - PoolId::new(binding.domain.id, ckf_dc_id), - EndpointLocator::new(ckf_dc_id, endpoint.clone()), - Some(EndpointLocator::new( - ckf_dc_id, - binding.kv_state_endpoint.clone(), - )), - ), - }; - let (handle, faults) = - KvDcRelayHandle::spawn_with_publication_delay(config, scope, publication.delay)?; +) -> anyhow::Result { + let attachment = pools + .attach(PoolAttachRequest { + pool_id: PoolId::new(binding.domain.id, ckf_dc_id), + endpoint: endpoint.clone(), + registrations: registrations.clone(), + }) + .await?; let initial_recoveries = membership_watch .borrow() .sources @@ -944,14 +1304,14 @@ async fn start_endpoint_actor( }) .collect(); let target = KvDcRelayRecoveryTarget::new( - handle.clone(), + attachment.handle.clone(), rebuild_permit, initial_recoveries, recovery_attempt_timeout, ); let recovery = match start_target_subscriber( - component, - endpoint, + component.clone(), + endpoint.clone(), target, membership_watch, "kv-dc-relay".to_string(), @@ -964,37 +1324,159 @@ async fn start_endpoint_actor( { Ok(recovery) => recovery, Err(error) => { - let _ = handle.shutdown().await; + // `detach` owns the actor fault receiver and is cancellation-sensitive; this + // endpoint-slot task is joined by the host, so keep the await inline and drive it to + // completion before returning. + let _ = pools.detach(attachment).await; return Err(error); } }; - Ok(EndpointActorRuntime { - handle, + Ok(EndpointPoolRuntime { + attachment, recovery, - faults, binding, + registrations, }) } -async fn stop_endpoint_actor(active: EndpointActorRuntime) { - active.recovery.shutdown().await; - if let Err(error) = active.handle.shutdown().await { - tracing::warn!(%error, endpoint = %active.handle.scope.pool_binding.serving_endpoint().endpoint_id(), "Failed to drain KV DC Relay endpoint actor"); +async fn stop_endpoint_pool(active: EndpointPoolRuntime, pools: &PoolRegistry) { + retire_endpoint_pool(active, pools, PoolRetirementMode::Graceful).await; +} + +async fn fence_endpoint_pool(active: EndpointPoolRuntime, pools: &PoolRegistry) { + retire_endpoint_pool(active, pools, PoolRetirementMode::Fenced).await; +} + +async fn retire_endpoint_pool( + active: EndpointPoolRuntime, + pools: &PoolRegistry, + mode: PoolRetirementMode, +) { + let EndpointPoolRuntime { + attachment, + recovery, + binding, + .. + } = active; + let PoolAttachment { + pool_id, + layout_generation, + handle, + mut faults, + .. + } = attachment; + let teardown = async { + match mode { + PoolRetirementMode::Graceful => { + recovery.shutdown().await; + handle.shutdown().await + } + PoolRetirementMode::Fenced => { + let ((), result) = tokio::join!(recovery.shutdown(), handle.fence()); + result + } + } + }; + let result = withdraw_drain_and_remove_pool( + pools, + pool_id, + layout_generation, + mode, + &mut faults, + teardown, + ) + .await; + if let Err(error) = result { + tracing::warn!( + %error, + %pool_id, + endpoint = %binding.kv_state_endpoint, + ?mode, + "Failed to retire KV DC Relay pool actor" + ); + } +} + +async fn withdraw_drain_and_remove_pool( + pools: &PoolRegistry, + pool_id: PoolId, + layout_generation: u64, + mode: PoolRetirementMode, + faults: &mut tokio::sync::mpsc::Receiver, + teardown: impl Future, +) -> T { + if !pools.withdraw(pool_id, layout_generation, mode).await { + tracing::warn!( + %pool_id, + layout_generation, + ?mode, + "KV DC Relay pool generation was already absent during retirement" + ); } + let result = drain_faults_while(pool_id, faults, teardown).await; + pools.remove(pool_id, layout_generation).await; + result } -async fn fence_endpoint_actor(active: EndpointActorRuntime) { - // Stop publication first. Recovery shutdown may attempt rank resets, but a producer whose - // exact state is suspect must not emit another apparently valid delta while being retired. - if let Err(error) = active.handle.fence().await { - tracing::warn!(%error, endpoint = %active.handle.scope.pool_binding.serving_endpoint().endpoint_id(), "Failed to fence KV DC Relay endpoint actor cleanly"); +enum FaultCollection { + Completed(T), + /// The source-local future was dropped; the caller must retire the producer generation. + ProducerFence(ProducerFenceTrigger), +} + +async fn collect_pending_while( + receiver: &mut tokio::sync::mpsc::Receiver, + pending: &mut PendingActorFaults, + future: impl Future, +) -> FaultCollection { + tokio::pin!(future); + loop { + tokio::select! { + result = &mut future => return FaultCollection::Completed(result), + item = receiver.recv() => match item { + Some(fault) => { + pending.push(fault); + if let Some(trigger) = pending.take_producer_fence() { + return FaultCollection::ProducerFence(trigger); + } + } + None => return FaultCollection::Completed(future.await), + }, + } } - active.recovery.shutdown().await; +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RegistrationRefresh { + Skipped, + Published, +} + +async fn refresh_pool_registrations( + pools: &PoolRegistry, + attachment: &mut PoolAttachment, + active_binding: &ActorBinding, + desired_binding: Option<&ActorBinding>, + source_binding_pending: bool, + desired_registrations: &[CanonicalModelRegistration], +) -> anyhow::Result { + if source_binding_pending || Some(active_binding) != desired_binding { + return Ok(RegistrationRefresh::Skipped); + } + + pools + .replace_registrations(attachment, desired_registrations.to_vec()) + .await?; + Ok(RegistrationRefresh::Published) +} + +fn target_fault_retirement_mode(disposition: TargetFaultDisposition) -> Option { + (disposition == TargetFaultDisposition::Fenced).then_some(PoolRetirementMode::Fenced) } #[cfg(feature = "ckf-diagnostics")] async fn endpoint_stats( - endpoint: EndpointId, + slot_id: EndpointId, status: SharedEndpointStatus, ) -> Result { let status = status.read().await.clone(); @@ -1008,9 +1490,9 @@ async fn endpoint_stats( Some(KvDcRelayAggregationStats { members: members .into_iter() - .map(|(worker, blocks)| KvDcRelayMemberStats { - worker_id: worker.worker_id, - dp_rank: worker.dp_rank, + .map(|(source, blocks)| KvDcRelayMemberStats { + worker_id: source.worker_id, + dp_rank: source.dp_rank, blocks, }) .collect(), @@ -1055,16 +1537,23 @@ async fn endpoint_stats( }; let membership = status.membership; Ok(KvDcRelayEndpointStats { - serving_endpoint: endpoint.to_string(), + serving_endpoint: slot_id.to_string(), lifecycle: status.lifecycle.as_str().to_string(), layout_generation: status.layout_generation, cache_domain: membership .as_ref() .and_then(|membership| membership.domain.as_ref()) .map(cache_domain_stats), - compatibility_conflict: membership + membership_conflicts: membership .as_ref() - .is_some_and(|membership| membership.compatibility_conflict), + .map(|membership| { + membership + .conflicts + .iter() + .map(|conflict| format!("{conflict:?}")) + .collect() + }) + .unwrap_or_default(), models: membership .as_ref() .map(|membership| membership.models.clone()) @@ -1162,53 +1651,768 @@ fn actor_health(handle: &KvDcRelayHandle) -> KvDcRelayActorStats { #[cfg(test)] mod tests { - use dynamo_kv_router::identity::{ - CacheSemanticsId, DcId, IdentitySource, IndexerDomainId, RoutingScopeId, + use dynamo_kv_router::{ + identity::{CacheSemanticsId, DcId, IdentitySource, IndexerDomainId, RoutingScopeId}, + indexer::cuckoo::{CkfFailureDisposition, CkfFailurePoint}, + protocols::{KV_EVENT_SUBJECT, WorkerWithDpRank}, + }; + use dynamo_runtime::{ + DistributedRuntime, Runtime, + discovery::{DiscoveryInstance, DiscoverySpec, EventTransportKind}, + distributed::DistributedConfig, + transports::event_plane::{EventPublisher, EventScope}, }; + use super::super::actor::ActorFaultCategory; use super::*; - use crate::kv_dc_relay::resolution::ResolvedIndexerDomain; + use crate::kv_router::indexer::SourceEpoch; - fn membership(endpoint: &str, domain: ResolvedIndexerDomain) -> EndpointMembership { + fn membership(endpoint: &str, domain: KvCacheDomainKey) -> EndpointMembership { + let endpoint = EndpointId::from(endpoint); EndpointMembership { - endpoint: EndpointId::from(endpoint), + endpoint, generation: 1, domain: Some(domain), - compatibility_conflict: false, - models: Vec::new(), + registrations: vec![CanonicalModelRegistration::new( + super::super::identity::CanonicalModelId::new("llama").unwrap(), + Vec::new(), + )], + models: vec!["llama".to_string()], aliases: Vec::new(), roles: Vec::new(), runtime_configs: HashMap::new(), + conflicts: Vec::new(), } } - #[test] - fn simultaneous_endpoints_cannot_own_one_pool() { - let domain_id = IndexerDomainId::new( - CacheSemanticsId::new([1; 16], IdentitySource::Explicit), - RoutingScopeId::new([2; 16], IdentitySource::Explicit), - ); - let domain = ResolvedIndexerDomain { - id: domain_id, - diagnostic_model_artifact: "model".to_string(), - kv_block_size: 512, + fn domain(seed: u8, artifact: &str) -> KvCacheDomainKey { + KvCacheDomainKey { + id: IndexerDomainId::new( + CacheSemanticsId::new([seed; 16], IdentitySource::Explicit), + RoutingScopeId::new([seed.wrapping_add(1); 16], IdentitySource::Explicit), + ), + diagnostic_model_artifact: artifact.to_string(), + kv_block_size: 64, event_hash_format: 1, + } + } + + fn registry() -> PoolRegistry { + PoolRegistry::new( + DcRelayIdentity::new(11, 7), + PoolActorConfig { + expected_unique_blocks: 32, + publication_threshold: 1, + publication_delay: Duration::from_millis(1), + }, + ) + } + + fn actor_fault( + worker_id: WorkerId, + dp_rank: DpRank, + source_epoch: u64, + event_id: u64, + disposition: CkfFailureDisposition, + ) -> ActorFault { + let category = match disposition.action { + CkfFailureAction::ReportResourceFailure => ActorFaultCategory::Resource, + CkfFailureAction::RejectSource => ActorFaultCategory::SourceProtocol, + CkfFailureAction::ContinueCapacityOmission + | CkfFailureAction::FenceAndRebuildProducer => ActorFaultCategory::ProducerInvariant, + CkfFailureAction::DeactivateAndSnapshot | CkfFailureAction::RetrySnapshot => { + unreachable!("consumer fault cannot be used by Relay host tests") + } }; - let first = membership("ns/router/first", domain.clone()); - let second = membership("ns/router/second", domain); + ActorFault { + worker_id, + dp_rank, + source_epoch: SourceEpoch::new(source_epoch), + event_id: Some(event_id), + category, + disposition, + message: format!("fault {event_id}"), + } + } + + async fn test_component(name: &str) -> Component { + let runtime = Runtime::from_current().unwrap(); + let drt = DistributedRuntime::new(runtime, DistributedConfig::process_local()) + .await + .unwrap(); + drt.namespace(format!("kv-dc-relay-{name}")) + .unwrap() + .component("relay") + .unwrap() + } + + async fn register_live_source( + component: &Component, + endpoint: &EndpointId, + worker: WorkerWithDpRank, + ) -> EventPublisher { + let publisher = EventPublisher::for_endpoint_id_with_transport( + component.drt(), + endpoint, + KV_EVENT_SUBJECT, + EventTransportKind::Zmq, + ) + .await + .unwrap(); + let source = crate::discovery::KvEventSource { + kv_state_endpoint: endpoint.clone(), + worker, + publisher_id: publisher.publisher_id(), + recovery_target: None, + }; + component + .drt() + .discovery() + .register(DiscoverySpec::EventSource { + scope: EventScope::Endpoint { + endpoint: endpoint.clone(), + }, + topic: KV_EVENT_SUBJECT.to_string(), + publisher_id: publisher.publisher_id(), + metadata: serde_json::to_value(source).unwrap(), + }) + .await + .unwrap(); + publisher + } + + fn projected_membership( + endpoint: &EndpointId, + worker_id: WorkerId, + model: &str, + artifact: &str, + kv_state_endpoint: EndpointId, + ) -> EndpointMembership { + projected_membership_with_metadata( + endpoint, + worker_id, + model, + artifact, + kv_state_endpoint, + Vec::new(), + None, + ) + } + + fn projected_membership_with_metadata( + endpoint: &EndpointId, + worker_id: WorkerId, + model: &str, + artifact: &str, + kv_state_endpoint: EndpointId, + aliases: Vec, + context_length: Option, + ) -> EndpointMembership { + let mut card = crate::model_card::ModelDeploymentCard::with_name_only(model); + card.source_path = Some(artifact.to_string()); + card.kv_cache_block_size = 64; + card.aliases = aliases; + card.runtime_config = ModelRuntimeConfig { + context_length, + data_parallel_start_rank: 0, + data_parallel_size: 1, + enable_local_indexer: true, + kv_state_endpoint: Some(kv_state_endpoint), + ..ModelRuntimeConfig::default() + }; + let instance = DiscoveryInstance::Model { + namespace: endpoint.namespace.clone(), + component: endpoint.component.clone(), + endpoint: endpoint.name.clone(), + instance_id: worker_id, + card_json: serde_json::to_value(card).unwrap(), + model_suffix: None, + }; + super::super::discovery::project_instances_for_test(vec![instance]) + .endpoints + .get(endpoint) + .cloned() + .unwrap() + } + + async fn wait_for_catalog( + receiver: &mut watch::Receiver, + predicate: impl Fn(&DcPoolCatalog) -> bool, + ) -> DcPoolCatalog { + tokio::time::timeout(Duration::from_secs(5), async { + loop { + let catalog = receiver.borrow_and_update().clone(); + if predicate(&catalog) { + return catalog; + } + receiver.changed().await.unwrap(); + } + }) + .await + .expect("Relay catalog transition timed out") + } + + #[tokio::test] + async fn departed_endpoint_slots_are_reaped_instead_of_parked() { + let slot_id = EndpointId::from("prod.backend.generate"); + let (metadata, mut metadata_rx) = watch::channel(None); + let task = tokio::spawn(async move { while metadata_rx.changed().await.is_ok() {} }); + let mut slots = HashMap::from([( + slot_id.clone(), + EndpointSlotTask { + metadata, + status: Arc::new(RwLock::new(EndpointSlotStatus::default())), + task, + }, + )]); + let mut retired_slots = JoinSet::new(); + + retire_departed_endpoint_slots( + &DcMembershipView::default(), + &mut slots, + &mut retired_slots, + ); + + assert!(slots.is_empty()); + let (retired_slot, result) = retired_slots.join_next().await.unwrap().unwrap(); + assert_eq!(retired_slot, slot_id); + result.unwrap(); + } + + #[tokio::test] + async fn repeated_relay_start_separates_drt_identity_from_relay_incarnation() { + let component = test_component("incarnation").await; + let first = KvDcRelay::start( + component.clone(), + "test-dc".to_string(), + KvDcRelayConfig::default(), + ) + .await + .unwrap(); + let first_identity = first.pool_catalog().identity(); + first.shutdown().await.unwrap(); + + let second = KvDcRelay::start(component, "test-dc".to_string(), KvDcRelayConfig::default()) + .await + .unwrap(); + let second_identity = second.pool_catalog().identity(); + second.shutdown().await.unwrap(); + + assert_eq!( + first_identity.drt_instance_id(), + second_identity.drt_instance_id() + ); + assert_ne!( + first_identity.relay_incarnation(), + second_identity.relay_incarnation() + ); + } + + #[tokio::test] + async fn duplicate_pool_owners_are_all_fenced_and_make_health_unhealthy() { + let domain = domain(1, "meta/llama"); + let first = membership("prod.backend.fast", domain.clone()); + let second = membership("prod.backend.slow", domain); let mut view = DcMembershipView { - endpoints: HashMap::from([ + endpoints: Arc::new(HashMap::from([ (first.endpoint.clone(), first), (second.endpoint.clone(), second), - ]), + ])), }; reject_duplicate_live_pools(&mut view, DcId::new(7)); + assert!(view.endpoints.values().all(|membership| { + !membership.is_materializable() + && inactive_slot_lifecycle(Some(membership)) == SlotLifecycle::Fenced + })); - assert!( + let statuses = Arc::new(RwLock::new( view.endpoints - .values() - .all(|membership| membership.compatibility_conflict) + .iter() + .map(|(endpoint, membership)| { + let status = EndpointSlotStatus { + lifecycle: inactive_slot_lifecycle(Some(membership)), + membership: Some(membership.clone()), + ..EndpointSlotStatus::default() + }; + (endpoint.clone(), Arc::new(RwLock::new(status))) + }) + .collect(), + )); + let relay = KvDcRelay { + #[cfg(feature = "ckf-diagnostics")] + dc_id: Arc::from("test-dc"), + #[cfg(feature = "ckf-diagnostics")] + relay_identity: DcRelayIdentity::new(11, 7), + cancel: CancellationToken::new(), + membership: Mutex::new(None), + supervisor: Mutex::new(None), + statuses, + pools: Arc::new(registry()), + }; + + let health = relay.health().await; + assert!(!health.healthy); + assert_eq!(health.fenced_endpoint_count, 2); + assert_eq!(health.active_endpoint_count, 0); + } + + #[tokio::test] + async fn binding_change_never_publishes_new_model_with_old_producer() { + let registry = registry(); + let mut catalog_rx = registry.watch_catalog(); + let old_domain = domain(1, "meta/llama"); + let new_domain = domain(3, "meta/llama-v2"); + let old_binding = ActorBinding { + domain: old_domain.clone(), + kv_state_endpoint: EndpointId::from("prod.backend.old-kv"), + }; + let desired_binding = ActorBinding { + domain: new_domain, + kv_state_endpoint: EndpointId::from("prod.backend.new-kv"), + }; + let mut old_attachment = registry + .attach(PoolAttachRequest { + pool_id: PoolId::new(old_domain.id, DcId::new(7)), + endpoint: EndpointId::from("prod.backend.generate"), + registrations: vec![CanonicalModelRegistration::new( + super::super::identity::CanonicalModelId::new("llama").unwrap(), + Vec::new(), + )], + }) + .await + .unwrap(); + catalog_rx.changed().await.unwrap(); + let old_catalog = catalog_rx.borrow_and_update().clone(); + let old_producer = old_catalog.pools()[0].producer(); + let new_registrations = vec![CanonicalModelRegistration::new( + super::super::identity::CanonicalModelId::new("llama-v2").unwrap(), + Vec::new(), + )]; + + assert_eq!( + refresh_pool_registrations( + ®istry, + &mut old_attachment, + &old_binding, + Some(&desired_binding), + false, + &new_registrations, + ) + .await + .unwrap(), + RegistrationRefresh::Skipped + ); + assert!(!catalog_rx.has_changed().unwrap()); + assert_eq!( + old_catalog.pools()[0].registrations()[0].model().as_str(), + "llama" + ); + + registry.detach(old_attachment).await.unwrap(); + catalog_rx.changed().await.unwrap(); + let withdrawn_catalog = catalog_rx.borrow_and_update().clone(); + assert!(withdrawn_catalog.pools().is_empty()); + + let new_attachment = registry + .attach(PoolAttachRequest { + pool_id: PoolId::new(desired_binding.domain.id, DcId::new(7)), + endpoint: EndpointId::from("prod.backend.generate"), + registrations: new_registrations, + }) + .await + .unwrap(); + catalog_rx.changed().await.unwrap(); + let new_catalog = catalog_rx.borrow_and_update().clone(); + assert_ne!(new_catalog.pools()[0].producer(), old_producer); + assert_eq!( + new_catalog.pools()[0].registrations()[0].model().as_str(), + "llama-v2" + ); + + for catalog in [&old_catalog, &withdrawn_catalog, &new_catalog] { + assert!(!catalog.pools().iter().any(|descriptor| { + descriptor.producer() == old_producer + && descriptor + .registrations() + .iter() + .any(|registration| registration.model().as_str() == "llama-v2") + })); + } + + registry.detach(new_attachment).await.unwrap(); + } + + #[tokio::test] + async fn mdc_binding_transition_never_publishes_new_model_with_old_producer() { + let component = test_component("mdc-transition").await; + let worker_id = component.drt().connection_id(); + let worker = WorkerWithDpRank::new(worker_id, 0); + let serving_endpoint = EndpointId::from("relay-test.backend.generate"); + let old_kv_endpoint = EndpointId::from("relay-test.backend.kv-old"); + let new_kv_endpoint = EndpointId::from("relay-test.backend.kv-new"); + let _old_publisher = register_live_source(&component, &old_kv_endpoint, worker).await; + let _new_publisher = register_live_source(&component, &new_kv_endpoint, worker).await; + let old_membership = projected_membership( + &serving_endpoint, + worker_id, + "llama", + "meta/llama", + old_kv_endpoint, + ); + let new_membership = projected_membership( + &serving_endpoint, + worker_id, + "llama-v2", + "meta/llama-v2", + new_kv_endpoint, + ); + let (metadata_tx, metadata_rx) = watch::channel(Some(old_membership)); + let status = Arc::new(RwLock::new(EndpointSlotStatus::default())); + let registry = Arc::new(registry()); + let mut catalog_rx = registry.watch_catalog(); + let slot_cancel = CancellationToken::new(); + let slot = tokio::spawn(run_endpoint_slot( + component, + DcId::new(7), + serving_endpoint, + metadata_rx, + status, + Arc::new(Semaphore::new(1)), + Arc::new(Semaphore::new(1)), + registry, + Duration::from_secs(1), + slot_cancel.clone(), + )); + + let old_catalog = wait_for_catalog(&mut catalog_rx, |catalog| { + catalog + .pools() + .iter() + .any(|descriptor| descriptor.registrations()[0].model().as_str() == "llama") + }) + .await; + let old_producer = old_catalog.pools()[0].producer(); + + metadata_tx.send_replace(Some(new_membership)); + let new_catalog = + tokio::time::timeout(Duration::from_secs(5), async { + loop { + catalog_rx.changed().await.unwrap(); + let catalog = catalog_rx.borrow_and_update().clone(); + assert!(!catalog.pools().iter().any(|descriptor| { + descriptor.producer() == old_producer + && descriptor + .registrations() + .iter() + .any(|registration| registration.model().as_str() == "llama-v2") + })); + if catalog.pools().iter().any(|descriptor| { + descriptor.registrations()[0].model().as_str() == "llama-v2" + }) { + return catalog; + } + } + }) + .await + .expect("replacement Relay catalog generation timed out"); + assert_ne!(new_catalog.pools()[0].producer(), old_producer); + + slot_cancel.cancel(); + tokio::time::timeout(Duration::from_secs(5), slot) + .await + .expect("endpoint slot shutdown timed out") + .unwrap(); + } + + #[tokio::test] + async fn alias_and_context_change_refreshes_the_active_pool_catalog() { + let component = test_component("metadata-transition").await; + let worker_id = component.drt().connection_id(); + let worker = WorkerWithDpRank::new(worker_id, 0); + let serving_endpoint = EndpointId::from("relay-test.backend.generate"); + let kv_endpoint = EndpointId::from("relay-test.backend.kv"); + let _publisher = register_live_source(&component, &kv_endpoint, worker).await; + let old_membership = projected_membership_with_metadata( + &serving_endpoint, + worker_id, + "llama", + "meta/llama", + kv_endpoint.clone(), + vec!["old-alias".to_string()], + Some(4096), + ); + let new_membership = projected_membership_with_metadata( + &serving_endpoint, + worker_id, + "llama", + "meta/llama", + kv_endpoint, + vec!["new-alias".to_string()], + Some(8192), + ); + let (metadata_tx, metadata_rx) = watch::channel(Some(old_membership)); + let status = Arc::new(RwLock::new(EndpointSlotStatus::default())); + let registry = Arc::new(registry()); + let mut catalog_rx = registry.watch_catalog(); + let slot_cancel = CancellationToken::new(); + let slot = tokio::spawn(run_endpoint_slot( + component, + DcId::new(7), + serving_endpoint, + metadata_rx, + status, + Arc::new(Semaphore::new(1)), + Arc::new(Semaphore::new(1)), + registry, + Duration::from_secs(1), + slot_cancel.clone(), + )); + + let old_catalog = wait_for_catalog(&mut catalog_rx, |catalog| { + catalog.pools().iter().any(|descriptor| { + descriptor.registrations()[0] + .aliases() + .iter() + .any(|alias| alias.as_str() == "old-alias") + }) + }) + .await; + let producer = old_catalog.pools()[0].producer(); + + metadata_tx.send_replace(Some(new_membership)); + let new_catalog = wait_for_catalog(&mut catalog_rx, |catalog| { + catalog.pools().iter().any(|descriptor| { + descriptor.producer() == producer + && descriptor.registrations()[0] + .aliases() + .iter() + .any(|alias| alias.as_str() == "new-alias") + }) + }) + .await; + assert!(new_catalog.pools().iter().all(|descriptor| { + descriptor.registrations().iter().all(|registration| { + registration + .aliases() + .iter() + .all(|alias| alias.as_str() != "old-alias") + }) + })); + + slot_cancel.cancel(); + tokio::time::timeout(Duration::from_secs(5), slot) + .await + .expect("endpoint slot shutdown timed out") + .unwrap(); + } + + #[tokio::test] + async fn pool_is_withdrawn_before_recovery_teardown_completes() { + let registry = Arc::new(registry()); + let attachment = registry + .attach(PoolAttachRequest { + pool_id: PoolId::new(domain(1, "meta/llama").id, DcId::new(7)), + endpoint: EndpointId::from("prod.backend.generate"), + registrations: vec![CanonicalModelRegistration::new( + super::super::identity::CanonicalModelId::new("llama").unwrap(), + Vec::new(), + )], + }) + .await + .unwrap(); + let PoolAttachment { + pool_id, + layout_generation, + handle, + mut faults, + .. + } = attachment; + let (teardown_started, teardown_started_rx) = tokio::sync::oneshot::channel(); + let (release_teardown, release_teardown_rx) = tokio::sync::oneshot::channel(); + let retirement_registry = registry.clone(); + let retirement = tokio::spawn(async move { + withdraw_drain_and_remove_pool( + &retirement_registry, + pool_id, + layout_generation, + PoolRetirementMode::Graceful, + &mut faults, + async move { + let _ = teardown_started.send(()); + let _ = release_teardown_rx.await; + handle.shutdown().await + }, + ) + .await + }); + + teardown_started_rx.await.unwrap(); + assert!(registry.catalog().pools().is_empty()); + assert_eq!(registry.pool_count().await, 1); + + release_teardown.send(()).unwrap(); + retirement.await.unwrap().unwrap(); + assert_eq!(registry.pool_count().await, 0); + } + + #[test] + fn pending_faults_coalesce_duplicate_and_stronger_source_actions() { + let resource = CkfFailurePoint::PrecommitAllocationFailure.disposition(); + let reject = CkfFailurePoint::SourceProtocolFailure.disposition(); + let mut pending = PendingActorFaults::default(); + + for event_id in 0..1_000 { + pending.push(actor_fault(1, 0, 1, event_id, resource)); + } + pending.push(actor_fault(1, 0, 1, 1_000, reject)); + pending.push(actor_fault(1, 0, 1, 1_001, resource)); + + assert_eq!(pending.len(), 1); + let Some(PendingActorAction::Fault(fault)) = pending.pop_front() else { + panic!("strongest source fault was not retained"); + }; + assert_eq!(fault.disposition.action, CkfFailureAction::RejectSource); + assert!(pending.pop_front().is_none()); + + pending.push(actor_fault(1, 0, 1, 1_002, reject)); + pending.push(actor_fault(1, 0, 2, 1_003, resource)); + let Some(PendingActorAction::Fault(fault)) = pending.pop_front() else { + panic!("newer source epoch was not retained"); + }; + assert_eq!(fault.source_epoch, SourceEpoch::new(2)); + assert_eq!( + fault.disposition.action, + CkfFailureAction::ReportResourceFailure + ); + } + + #[test] + fn pending_fault_overflow_fails_safe_with_a_producer_fence() { + let resource = CkfFailurePoint::PrecommitAllocationFailure.disposition(); + let mut pending = PendingActorFaults::default(); + + for worker_id in 0..MAX_PENDING_SOURCE_FAULTS as u64 { + pending.push(actor_fault(worker_id, 0, 1, worker_id, resource)); + } + assert_eq!(pending.len(), MAX_PENDING_SOURCE_FAULTS); + + pending.push(actor_fault( + MAX_PENDING_SOURCE_FAULTS as u64, + 0, + 1, + MAX_PENDING_SOURCE_FAULTS as u64, + resource, + )); + assert_eq!(pending.len(), 1); + assert!(matches!( + pending.pop_front(), + Some(PendingActorAction::ProducerFence( + ProducerFenceTrigger::PendingOverflow(_) + )) + )); + assert!(pending.pop_front().is_none()); + } + + #[test] + fn producer_fence_supersedes_varied_pending_source_faults() { + let resource = CkfFailurePoint::PrecommitAllocationFailure.disposition(); + let reject = CkfFailurePoint::SourceProtocolFailure.disposition(); + let fence = CkfFailurePoint::PrewriteInvariantMismatch.disposition(); + let mut pending = PendingActorFaults::default(); + + pending.push(actor_fault(1, 0, 1, 1, resource)); + pending.push(actor_fault(1, 0, 1, 2, resource)); + pending.push(actor_fault(2, 0, 1, 3, reject)); + pending.push(actor_fault(3, 0, 1, 4, resource)); + pending.push(actor_fault(1, 0, 1, 5, fence)); + pending.push(actor_fault(4, 0, 1, 6, resource)); + + assert_eq!(pending.len(), 1); + let Some(PendingActorAction::ProducerFence(ProducerFenceTrigger::Fault(fault))) = + pending.pop_front() + else { + panic!("producer fence did not supersede weaker source faults"); + }; + assert_eq!( + fault.disposition.action, + CkfFailureAction::FenceAndRebuildProducer + ); + assert!(pending.pop_front().is_none()); + } + + #[tokio::test] + async fn producer_fence_interrupts_inflight_fault_recovery() { + let resource = CkfFailurePoint::PrecommitAllocationFailure.disposition(); + let fence = CkfFailurePoint::PrewriteInvariantMismatch.disposition(); + let (sender, mut receiver) = tokio::sync::mpsc::channel(DEFAULT_FAULT_CAPACITY); + let (recovery_started, recovery_started_rx) = tokio::sync::oneshot::channel(); + let (release_recovery, release_recovery_rx) = tokio::sync::oneshot::channel::<()>(); + let collection = tokio::spawn(async move { + let mut pending = PendingActorFaults::default(); + let outcome = collect_pending_while(&mut receiver, &mut pending, async move { + let _ = recovery_started.send(()); + let _ = release_recovery_rx.await; + TargetFaultDisposition::Recovering + }) + .await; + (outcome, pending) + }); + + recovery_started_rx.await.unwrap(); + for event_id in 0..100 { + sender + .send(actor_fault(1, 0, 1, event_id, resource)) + .await + .unwrap(); + } + sender.send(actor_fault(1, 0, 1, 100, fence)).await.unwrap(); + + let (outcome, pending) = tokio::time::timeout(Duration::from_secs(1), collection) + .await + .expect("producer fence did not interrupt fault recovery") + .unwrap(); + assert!(matches!( + outcome, + FaultCollection::ProducerFence(ProducerFenceTrigger::Fault(_)) + )); + assert_eq!(pending.len(), 0); + assert!(release_recovery.send(()).is_err()); + } + + #[tokio::test] + async fn fenced_target_disposition_withdraws_the_pool_generation() { + let registry = registry(); + let domain = domain(1, "meta/llama"); + let attachment = registry + .attach(PoolAttachRequest { + pool_id: PoolId::new(domain.id, DcId::new(7)), + endpoint: EndpointId::from("prod.backend.generate"), + registrations: vec![CanonicalModelRegistration::new( + super::super::identity::CanonicalModelId::new("llama").unwrap(), + Vec::new(), + )], + }) + .await + .unwrap(); + let mode = target_fault_retirement_mode(TargetFaultDisposition::Fenced).unwrap(); + assert!( + registry + .withdraw(attachment.pool_id, attachment.layout_generation, mode) + .await ); + assert!(registry.catalog().pools().is_empty()); + + let PoolAttachment { + pool_id, + layout_generation, + handle, + mut faults, + .. + } = attachment; + drain_faults_while(pool_id, &mut faults, handle.fence()) + .await + .unwrap(); + assert!(registry.remove(pool_id, layout_generation).await); } } diff --git a/lib/llm/src/kv_dc_relay/identity.rs b/lib/llm/src/kv_dc_relay/identity.rs new file mode 100644 index 000000000000..e92b147f4e7c --- /dev/null +++ b/lib/llm/src/kv_dc_relay/identity.rs @@ -0,0 +1,469 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::collections::BTreeMap; +use std::fmt; +use std::sync::{Arc, OnceLock}; + +use dynamo_kv_router::identity::{IdentitySource, PoolId}; +use dynamo_kv_router::indexer::cuckoo::ProducerIdentity; +use dynamo_runtime::protocols::EndpointId; +use serde::{Deserialize, Deserializer, Serialize}; + +fn validate_identity_text( + value: impl Into, + empty: E, + surrounding_whitespace: E, +) -> Result { + let value = value.into(); + if value.is_empty() { + return Err(empty); + } + if value.trim() != value { + return Err(surrounding_whitespace); + } + Ok(value) +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)] +#[serde(transparent)] +pub struct CanonicalModelId(String); + +impl CanonicalModelId { + pub fn new(value: impl Into) -> Result { + validate_identity_text( + value, + CanonicalModelIdError::Empty, + CanonicalModelIdError::SurroundingWhitespace, + ) + .map(Self) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Display for CanonicalModelId { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(&self.0) + } +} + +impl<'de> Deserialize<'de> for CanonicalModelId { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let value = String::deserialize(deserializer)?; + Self::new(value).map_err(serde::de::Error::custom) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] +pub enum CanonicalModelIdError { + #[error("canonical model ID must not be empty")] + Empty, + #[error("canonical model ID must not contain leading or trailing whitespace")] + SurroundingWhitespace, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)] +#[serde(transparent)] +pub struct ModelAlias(String); + +impl ModelAlias { + pub fn new(value: impl Into) -> Result { + validate_identity_text( + value, + ModelAliasError::Empty, + ModelAliasError::SurroundingWhitespace, + ) + .map(Self) + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Display for ModelAlias { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str(&self.0) + } +} + +impl<'de> Deserialize<'de> for ModelAlias { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let value = String::deserialize(deserializer)?; + Self::new(value).map_err(serde::de::Error::custom) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] +pub enum ModelAliasError { + #[error("model alias must not be empty")] + Empty, + #[error("model alias must not contain leading or trailing whitespace")] + SurroundingWhitespace, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize)] +pub struct CanonicalModelRegistration { + model: CanonicalModelId, + target: ModelTarget, + aliases: Vec, +} + +impl CanonicalModelRegistration { + pub fn new(model: CanonicalModelId, aliases: Vec) -> Self { + let target = ModelTarget::Base { + base_model: model.clone(), + }; + Self::with_target(model, target, aliases) + } + + pub fn with_target( + model: CanonicalModelId, + target: ModelTarget, + mut aliases: Vec, + ) -> Self { + aliases.retain(|alias| alias.as_str() != model.as_str()); + aliases.sort_unstable(); + aliases.dedup(); + Self { + model, + target, + aliases, + } + } + + pub const fn model(&self) -> &CanonicalModelId { + &self.model + } + + pub const fn target(&self) -> &ModelTarget { + &self.target + } + + pub fn aliases(&self) -> &[ModelAlias] { + &self.aliases + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)] +pub enum ModelTarget { + Base { + base_model: CanonicalModelId, + }, + Lora { + base_model: CanonicalModelId, + adapter: CanonicalModelId, + }, +} + +impl ModelTarget { + pub const fn base_model(&self) -> &CanonicalModelId { + match self { + Self::Base { base_model } | Self::Lora { base_model, .. } => base_model, + } + } + + pub const fn adapter(&self) -> Option<&CanonicalModelId> { + match self { + Self::Base { .. } => None, + Self::Lora { adapter, .. } => Some(adapter), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize)] +pub struct DcPoolDescriptor { + producer: ProducerIdentity, + serving_endpoint: EndpointId, + registrations: Arc<[CanonicalModelRegistration]>, +} + +impl DcPoolDescriptor { + pub(crate) fn new( + producer: ProducerIdentity, + serving_endpoint: EndpointId, + registrations: Arc<[CanonicalModelRegistration]>, + ) -> Self { + Self { + producer, + serving_endpoint, + registrations, + } + } + + pub const fn producer(&self) -> ProducerIdentity { + self.producer + } + + pub const fn pool_id(&self) -> PoolId { + self.producer.pool_id() + } + + pub const fn serving_endpoint(&self) -> &EndpointId { + &self.serving_endpoint + } + + pub fn registrations(&self) -> &[CanonicalModelRegistration] { + &self.registrations + } +} + +/// Identity of one DC Relay runtime. +/// +/// `drt_instance_id` identifies the backing Dynamo runtime and can remain stable across an +/// in-process Relay restart. `relay_incarnation` is generated for every [`KvDcRelay::start`] +/// and fences producer generations created by different Relay lifetimes. +/// +/// [`KvDcRelay::start`]: super::KvDcRelay::start +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize)] +pub struct DcRelayIdentity { + drt_instance_id: u64, + relay_incarnation: u64, +} + +impl DcRelayIdentity { + pub const fn new(drt_instance_id: u64, relay_incarnation: u64) -> Self { + Self { + drt_instance_id, + relay_incarnation, + } + } + + pub const fn drt_instance_id(self) -> u64 { + self.drt_instance_id + } + + pub const fn relay_incarnation(self) -> u64 { + self.relay_incarnation + } +} + +#[derive(Debug)] +struct DcPoolCatalogPools { + by_id: BTreeMap, + ordered: OnceLock>, +} + +impl Clone for DcPoolCatalogPools { + fn clone(&self) -> Self { + Self { + by_id: self.by_id.clone(), + ordered: OnceLock::new(), + } + } +} + +#[derive(Clone)] +pub struct DcPoolCatalog { + identity: DcRelayIdentity, + revision: u64, + pools: DcPoolCatalogPools, +} + +impl DcPoolCatalog { + pub(crate) fn new( + identity: DcRelayIdentity, + revision: u64, + pools: Vec, + ) -> Self { + let by_id = pools + .into_iter() + .map(|descriptor| (descriptor.pool_id(), descriptor)) + .collect(); + Self { + identity, + revision, + pools: DcPoolCatalogPools { + by_id, + ordered: OnceLock::new(), + }, + } + } + + pub(crate) fn upsert(&mut self, revision: u64, descriptor: DcPoolDescriptor) { + self.revision = revision; + self.pools.by_id.insert(descriptor.pool_id(), descriptor); + self.pools.ordered.take(); + } + + pub(crate) fn remove(&mut self, revision: u64, pool_id: PoolId) { + self.revision = revision; + self.pools.by_id.remove(&pool_id); + self.pools.ordered.take(); + } + + pub(crate) fn clear(&mut self, revision: u64) { + self.revision = revision; + self.pools.by_id.clear(); + self.pools.ordered.take(); + } + + pub const fn identity(&self) -> DcRelayIdentity { + self.identity + } + + pub const fn drt_instance_id(&self) -> u64 { + self.identity.drt_instance_id() + } + + pub const fn relay_incarnation(&self) -> u64 { + self.identity.relay_incarnation() + } + + pub const fn revision(&self) -> u64 { + self.revision + } + + pub fn pools(&self) -> &[DcPoolDescriptor] { + self.pools + .ordered + .get_or_init(|| self.pools.by_id.values().cloned().collect()) + } + + #[cfg(test)] + pub(crate) fn is_materialized(&self) -> bool { + self.pools.ordered.get().is_some() + } +} + +impl fmt::Debug for DcPoolCatalog { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("DcPoolCatalog") + .field("identity", &self.identity) + .field("revision", &self.revision) + .field("pools", &self.pools()) + .finish() + } +} + +impl PartialEq for DcPoolCatalog { + fn eq(&self, other: &Self) -> bool { + self.identity == other.identity + && self.revision == other.revision + && self.pools.by_id == other.pools.by_id + } +} + +impl Eq for DcPoolCatalog {} + +impl Serialize for DcPoolCatalog { + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + #[derive(Serialize)] + struct Catalog<'a> { + drt_instance_id: u64, + relay_incarnation: u64, + revision: u64, + pools: &'a [DcPoolDescriptor], + } + + Catalog { + drt_instance_id: self.identity.drt_instance_id(), + relay_incarnation: self.identity.relay_incarnation(), + revision: self.revision, + pools: self.pools(), + } + .serialize(serializer) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub struct PoolIdentitySources { + cache_semantics: IdentitySource, + routing_scope: IdentitySource, +} + +impl PoolIdentitySources { + pub const fn from_pool(pool_id: PoolId) -> Self { + Self { + cache_semantics: pool_id.indexer_domain().cache_semantics().source(), + routing_scope: pool_id.indexer_domain().routing_scope().source(), + } + } + + pub const fn cache_semantics(self) -> IdentitySource { + self.cache_semantics + } + + pub const fn routing_scope(self) -> IdentitySource { + self.routing_scope + } + + pub const fn is_derived(self) -> bool { + matches!(self.cache_semantics, IdentitySource::DefaultDerived) + || matches!(self.routing_scope, IdentitySource::DefaultDerived) + } +} + +#[cfg(test)] +mod tests { + use dynamo_kv_router::identity::{CacheSemanticsId, DcId, IndexerDomainId, RoutingScopeId}; + + use super::*; + + fn pool(cache_source: IdentitySource, routing_source: IdentitySource) -> PoolId { + PoolId::new( + IndexerDomainId::new( + CacheSemanticsId::new([1; 16], cache_source), + RoutingScopeId::new([2; 16], routing_source), + ), + DcId::new(3), + ) + } + + #[test] + fn canonical_model_id_rejects_ambiguous_text() { + assert_eq!(CanonicalModelId::new(""), Err(CanonicalModelIdError::Empty)); + assert_eq!( + CanonicalModelId::new(" llama"), + Err(CanonicalModelIdError::SurroundingWhitespace) + ); + } + + #[test] + fn canonical_registration_normalizes_aliases_without_creating_self_alias() { + let model = CanonicalModelId::new("llama").unwrap(); + let registration = CanonicalModelRegistration::new( + model.clone(), + vec![ + ModelAlias::new("chat").unwrap(), + ModelAlias::new("llama").unwrap(), + ModelAlias::new("chat").unwrap(), + ], + ); + + assert_eq!(registration.model(), &model); + assert_eq!(registration.aliases(), &[ModelAlias::new("chat").unwrap()]); + } + + #[test] + fn pool_identity_sources_report_derived_components() { + let explicit = PoolIdentitySources::from_pool(pool( + IdentitySource::Explicit, + IdentitySource::Explicit, + )); + assert_eq!(explicit.cache_semantics(), IdentitySource::Explicit); + assert_eq!(explicit.routing_scope(), IdentitySource::Explicit); + assert!(!explicit.is_derived()); + + let derived = PoolIdentitySources::from_pool(pool( + IdentitySource::Explicit, + IdentitySource::DefaultDerived, + )); + assert_eq!(derived.cache_semantics(), IdentitySource::Explicit); + assert_eq!(derived.routing_scope(), IdentitySource::DefaultDerived); + assert!(derived.is_derived()); + } +} diff --git a/lib/llm/src/kv_dc_relay/pool_registry.rs b/lib/llm/src/kv_dc_relay/pool_registry.rs new file mode 100644 index 000000000000..06dd908a11d9 --- /dev/null +++ b/lib/llm/src/kv_dc_relay/pool_registry.rs @@ -0,0 +1,1100 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::collections::HashMap; +use std::future::Future; +use std::sync::Arc; +use std::time::Duration; + +use dynamo_kv_router::identity::PoolId; +use dynamo_kv_router::indexer::cuckoo::{CkfBuildError, CkfConfig, DcCkfState, ProducerIdentity}; +use dynamo_runtime::protocols::EndpointId; +use parking_lot::Mutex; +use tokio::sync::{Semaphore, mpsc, watch}; +use tokio_util::sync::CancellationToken; + +use super::actor::{ActorFault, KvDcRelayHandle, StreamScope}; +use super::host::KvDcRelayError; +use super::identity::{ + CanonicalModelId, CanonicalModelRegistration, DcPoolCatalog, DcPoolDescriptor, DcRelayIdentity, + ModelAlias, +}; + +const DEFAULT_CKF_ALLOCATION_CONCURRENCY: usize = 2; + +#[derive(Debug, Clone, Copy)] +pub(super) struct PoolActorConfig { + pub(super) expected_unique_blocks: usize, + pub(super) publication_threshold: usize, + pub(super) publication_delay: Duration, +} + +#[derive(Debug)] +pub(super) struct PoolAttachRequest { + pub(super) pool_id: PoolId, + pub(super) endpoint: EndpointId, + pub(super) registrations: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum PoolRetirementMode { + Graceful, + Fenced, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PoolEntryState { + Active, + Withdrawn, + Fenced, +} + +struct PoolEntry { + endpoint: EndpointId, + handle: KvDcRelayHandle, + identity: ProducerIdentity, + layout_generation: u64, + registrations: Arc<[CanonicalModelRegistration]>, + cancel: CancellationToken, + state: PoolEntryState, +} + +struct PoolReservation { + endpoint: EndpointId, + layout_generation: u64, +} + +struct PoolReservationGuard<'a> { + state: &'a Mutex, + pool_id: PoolId, + layout_generation: u64, + is_armed: bool, +} + +impl<'a> PoolReservationGuard<'a> { + fn new(state: &'a Mutex, pool_id: PoolId, layout_generation: u64) -> Self { + Self { + state, + pool_id, + layout_generation, + is_armed: true, + } + } + + fn disarm(&mut self) { + self.is_armed = false; + } +} + +impl Drop for PoolReservationGuard<'_> { + fn drop(&mut self) { + if self.is_armed { + rollback_reservation(&mut self.state.lock(), self.pool_id, self.layout_generation); + } + } +} + +struct PoolRegistryState { + pools: HashMap, + reservations: HashMap, + next_layout_generation: u64, + catalog_revision: u64, + accepting: bool, +} + +impl Default for PoolRegistryState { + fn default() -> Self { + Self { + pools: HashMap::new(), + reservations: HashMap::new(), + next_layout_generation: 1, + catalog_revision: 0, + accepting: true, + } + } +} + +pub(super) struct PoolAttachment { + pub(super) pool_id: PoolId, + pub(super) layout_generation: u64, + pub(super) handle: KvDcRelayHandle, + registrations: Arc<[CanonicalModelRegistration]>, + pub(super) faults: mpsc::Receiver, + pub(super) pool_cancel: CancellationToken, +} + +pub(super) struct PoolRegistry { + relay_identity: DcRelayIdentity, + actor_config: PoolActorConfig, + ckf_allocation_permits: Arc, + state: Mutex, + catalog_tx: watch::Sender, +} + +impl PoolRegistry { + pub(super) fn new(relay_identity: DcRelayIdentity, actor_config: PoolActorConfig) -> Self { + let (catalog_tx, _) = watch::channel(DcPoolCatalog::new(relay_identity, 0, Vec::new())); + Self { + relay_identity, + actor_config, + ckf_allocation_permits: Arc::new(Semaphore::new(DEFAULT_CKF_ALLOCATION_CONCURRENCY)), + state: Mutex::new(PoolRegistryState::default()), + catalog_tx, + } + } + + pub(super) async fn attach( + &self, + request: PoolAttachRequest, + ) -> anyhow::Result { + self.attach_with_builder(request, DcCkfState::new).await + } + + async fn attach_with_builder( + &self, + request: PoolAttachRequest, + builder: Builder, + ) -> anyhow::Result + where + Builder: FnOnce(CkfConfig) -> Result + Send + 'static, + { + anyhow::ensure!( + !request.registrations.is_empty(), + "pool {} requires at least one canonical model binding", + request.pool_id + ); + validate_registrations(&request.registrations)?; + + let layout_generation = { + let mut state = self.state.lock(); + anyhow::ensure!( + state.accepting, + "KV DC Relay pool registry is shutting down" + ); + if let Some(endpoint) = pool_owner(&state, request.pool_id) { + anyhow::bail!( + "pool {} is already owned by endpoint {} and cannot also attach endpoint {}", + request.pool_id, + endpoint, + request.endpoint + ); + } + let layout_generation = allocate_layout_generation(&mut state)?; + state.reservations.insert( + request.pool_id, + PoolReservation { + endpoint: request.endpoint.clone(), + layout_generation, + }, + ); + layout_generation + }; + let mut reservation = + PoolReservationGuard::new(&self.state, request.pool_id, layout_generation); + + let mut config = CkfConfig::new(self.actor_config.expected_unique_blocks); + config.publish_every_n_events = self.actor_config.publication_threshold; + let permit = self + .ckf_allocation_permits + .clone() + .acquire_owned() + .await + .map_err(|_| anyhow::anyhow!("KV DC Relay pool registry is shutting down"))?; + let ckf_state = tokio::task::spawn_blocking(move || { + let result = builder(config); + drop(permit); + result + }) + .await + .map_err(|error| anyhow::anyhow!("KV DC Relay CKF allocation task failed: {error}"))??; + + let registrations: Arc<[CanonicalModelRegistration]> = request.registrations.into(); + let cancel = CancellationToken::new(); + + let mut state = self.state.lock(); + let reservation_matches = state + .reservations + .get(&request.pool_id) + .is_some_and(|reservation| reservation.layout_generation == layout_generation); + anyhow::ensure!( + state.accepting && reservation_matches, + "pool {} generation {} reservation was retired before commit", + request.pool_id, + layout_generation + ); + let (handle, faults) = KvDcRelayHandle::spawn_with_state_and_publication_delay( + ckf_state, + StreamScope { + relay_incarnation: self.relay_identity.relay_incarnation(), + layout_generation, + pool_id: request.pool_id, + }, + self.actor_config.publication_delay, + ); + let identity = handle.identity(); + let descriptor = + DcPoolDescriptor::new(identity, request.endpoint.clone(), registrations.clone()); + state.reservations.remove(&request.pool_id); + debug_assert!(!state.pools.contains_key(&request.pool_id)); + state.pools.insert( + request.pool_id, + PoolEntry { + endpoint: request.endpoint.clone(), + handle: handle.clone(), + identity, + layout_generation, + registrations: registrations.clone(), + cancel: cancel.clone(), + state: PoolEntryState::Active, + }, + ); + publish_catalog_upsert(&mut state, &self.catalog_tx, descriptor); + reservation.disarm(); + + Ok(PoolAttachment { + pool_id: request.pool_id, + layout_generation, + handle, + registrations, + faults, + pool_cancel: cancel, + }) + } + + pub(super) async fn detach(&self, attachment: PoolAttachment) -> Result<(), KvDcRelayError> { + let PoolAttachment { + pool_id, + layout_generation, + handle, + mut faults, + .. + } = attachment; + self.withdraw(pool_id, layout_generation, PoolRetirementMode::Graceful) + .await; + let result = drain_faults_while(pool_id, &mut faults, handle.shutdown()).await; + self.remove(pool_id, layout_generation).await; + result + } + + pub(super) async fn replace_registrations( + &self, + attachment: &mut PoolAttachment, + registrations: Vec, + ) -> anyhow::Result<()> { + anyhow::ensure!( + !registrations.is_empty(), + "pool {} requires at least one canonical model binding", + attachment.pool_id + ); + if attachment.registrations.as_ref() == registrations.as_slice() { + return Ok(()); + } + + validate_registrations(®istrations)?; + let registrations: Arc<[CanonicalModelRegistration]> = registrations.into(); + let mut state = self.state.lock(); + let entry = state + .pools + .get_mut(&attachment.pool_id) + .ok_or_else(|| anyhow::anyhow!("pool {} is not attached", attachment.pool_id))?; + anyhow::ensure!( + entry.layout_generation == attachment.layout_generation + && entry.state == PoolEntryState::Active, + "pool {} generation {} is no longer active", + attachment.pool_id, + attachment.layout_generation + ); + entry.registrations = registrations.clone(); + let descriptor = DcPoolDescriptor::new( + entry.identity, + entry.endpoint.clone(), + registrations.clone(), + ); + attachment.registrations = registrations; + publish_catalog_upsert(&mut state, &self.catalog_tx, descriptor); + Ok(()) + } + + pub(super) async fn withdraw( + &self, + pool_id: PoolId, + layout_generation: u64, + mode: PoolRetirementMode, + ) -> bool { + let mut state = self.state.lock(); + let Some(entry) = state.pools.get_mut(&pool_id) else { + return false; + }; + if entry.layout_generation != layout_generation { + return false; + } + let was_active = entry.state == PoolEntryState::Active; + entry.state = match (entry.state, mode) { + (PoolEntryState::Fenced, _) | (_, PoolRetirementMode::Fenced) => PoolEntryState::Fenced, + _ => PoolEntryState::Withdrawn, + }; + entry.cancel.cancel(); + if was_active { + publish_catalog_remove(&mut state, &self.catalog_tx, pool_id); + } + true + } + + pub(super) async fn remove(&self, pool_id: PoolId, layout_generation: u64) -> bool { + let mut state = self.state.lock(); + let Some(entry) = state.pools.get(&pool_id) else { + return false; + }; + if entry.layout_generation != layout_generation { + return false; + } + let was_active = entry.state == PoolEntryState::Active; + let Some(entry) = state.pools.remove(&pool_id) else { + return false; + }; + entry.cancel.cancel(); + if was_active { + publish_catalog_remove(&mut state, &self.catalog_tx, pool_id); + } + true + } + + pub(super) fn catalog(&self) -> DcPoolCatalog { + self.catalog_tx.borrow().clone() + } + + pub(super) fn watch_catalog(&self) -> watch::Receiver { + self.catalog_tx.subscribe() + } + + pub(super) async fn shutdown(&self) { + let entries = { + let mut state = self.state.lock(); + state.accepting = false; + self.ckf_allocation_permits.close(); + state.reservations.clear(); + let entries = state.pools.drain().collect::>(); + publish_catalog_clear(&mut state, &self.catalog_tx); + entries + }; + for (pool_id, entry) in entries { + entry.cancel.cancel(); + if let Err(error) = entry.handle.fence().await { + tracing::warn!(%pool_id, %error, "KV Relay pool actor failed to fence during registry shutdown"); + } + } + } + + #[cfg(test)] + pub(super) async fn pool_count(&self) -> usize { + self.state.lock().pools.len() + } +} + +pub(super) async fn drain_faults_while( + pool_id: PoolId, + faults: &mut mpsc::Receiver, + future: impl Future, +) -> T { + tokio::pin!(future); + loop { + tokio::select! { + result = &mut future => return result, + fault = faults.recv() => match fault { + Some(fault) => tracing::debug!( + %pool_id, + worker_id = fault.worker_id, + dp_rank = fault.dp_rank, + category = ?fault.category, + error = %fault.message, + "Draining KV DC Relay actor fault while retiring its pool" + ), + None => return future.await, + }, + } + } +} + +fn allocate_layout_generation(state: &mut PoolRegistryState) -> anyhow::Result { + let generation = state.next_layout_generation; + state.next_layout_generation = generation + .checked_add(1) + .ok_or_else(|| anyhow::anyhow!("KV DC Relay layout generation space exhausted"))?; + Ok(generation) +} + +fn pool_owner(state: &PoolRegistryState, pool_id: PoolId) -> Option<&EndpointId> { + state + .pools + .get(&pool_id) + .map(|entry| &entry.endpoint) + .or_else(|| { + state + .reservations + .get(&pool_id) + .map(|reservation| &reservation.endpoint) + }) +} + +fn rollback_reservation(state: &mut PoolRegistryState, pool_id: PoolId, layout_generation: u64) { + if state + .reservations + .get(&pool_id) + .is_some_and(|reservation| reservation.layout_generation == layout_generation) + { + state.reservations.remove(&pool_id); + } +} + +fn advance_catalog_revision(state: &mut PoolRegistryState) -> u64 { + state.catalog_revision = state.catalog_revision.saturating_add(1); + state.catalog_revision +} + +fn publish_catalog_upsert( + state: &mut PoolRegistryState, + sender: &watch::Sender, + descriptor: DcPoolDescriptor, +) { + let revision = advance_catalog_revision(state); + sender.send_modify(|catalog| catalog.upsert(revision, descriptor)); +} + +fn publish_catalog_remove( + state: &mut PoolRegistryState, + sender: &watch::Sender, + pool_id: PoolId, +) { + let revision = advance_catalog_revision(state); + sender.send_modify(|catalog| catalog.remove(revision, pool_id)); +} + +fn publish_catalog_clear(state: &mut PoolRegistryState, sender: &watch::Sender) { + let revision = advance_catalog_revision(state); + sender.send_modify(|catalog| catalog.clear(revision)); +} + +fn validate_registrations(registrations: &[CanonicalModelRegistration]) -> anyhow::Result<()> { + let mut request_models = HashMap::with_capacity(registrations.len()); + for registration in registrations { + if let Some(previous) = + request_models.insert(registration.model().clone(), registration.target().clone()) + { + anyhow::ensure!( + previous == *registration.target(), + "canonical model {} resolves to conflicting targets in the same pool", + registration.model() + ); + anyhow::bail!( + "duplicate canonical model registration {}", + registration.model() + ); + } + } + + let mut request_aliases = HashMap::::new(); + for registration in registrations { + for alias in registration.aliases() { + let alias_as_model = CanonicalModelId::new(alias.as_str().to_string())?; + anyhow::ensure!( + !request_models.contains_key(&alias_as_model), + "model alias {} conflicts with canonical model {} in the same pool", + alias, + alias_as_model + ); + if let Some(owner) = request_aliases.insert(alias.clone(), registration.model().clone()) + { + anyhow::ensure!( + owner == *registration.model(), + "model alias {} is claimed by both {} and {} in the same pool", + alias, + owner, + registration.model() + ); + } + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use std::sync::{Arc, mpsc as std_mpsc}; + + use dynamo_kv_router::identity::{ + CacheSemanticsId, DcId, IdentitySource, IndexerDomainId, RoutingScopeId, + }; + + use super::*; + use crate::kv_dc_relay::identity::ModelTarget; + + type TestCkfBuilder = + Box Result + Send + 'static>; + + struct GatedBuilder { + builder: TestCkfBuilder, + started: tokio::sync::oneshot::Receiver<()>, + release: std_mpsc::Sender<()>, + finished: tokio::sync::oneshot::Receiver<()>, + } + + fn pool(seed: u8) -> PoolId { + PoolId::new( + IndexerDomainId::new( + CacheSemanticsId::new([seed; 16], IdentitySource::Explicit), + RoutingScopeId::new([seed.wrapping_add(1); 16], IdentitySource::Explicit), + ), + DcId::new(3), + ) + } + + fn config() -> PoolActorConfig { + PoolActorConfig { + expected_unique_blocks: 32, + publication_threshold: 1, + publication_delay: Duration::from_millis(1), + } + } + + fn relay_identity() -> DcRelayIdentity { + DcRelayIdentity::new(11, 7) + } + + fn registration(model: &str) -> CanonicalModelRegistration { + CanonicalModelRegistration::new( + CanonicalModelId::new(model).unwrap(), + vec![ModelAlias::new(format!("{model}-alias")).unwrap()], + ) + } + + fn request(pool_id: PoolId, endpoint: &str, model: &str) -> PoolAttachRequest { + PoolAttachRequest { + pool_id, + endpoint: EndpointId::from(endpoint), + registrations: vec![registration(model)], + } + } + + fn gated_builder() -> GatedBuilder { + let (started_tx, started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = std_mpsc::channel(); + let (finished_tx, finished_rx) = tokio::sync::oneshot::channel(); + let builder = move |config| { + let _ = started_tx.send(()); + release_rx + .recv_timeout(Duration::from_secs(1)) + .expect("test must release the CKF builder"); + let result = DcCkfState::new(config); + let _ = finished_tx.send(()); + result + }; + GatedBuilder { + builder: Box::new(builder), + started: started_rx, + release: release_tx, + finished: finished_rx, + } + } + + fn descriptor(catalog: &DcPoolCatalog, pool_id: PoolId) -> &DcPoolDescriptor { + catalog + .pools() + .iter() + .find(|descriptor| descriptor.pool_id() == pool_id) + .unwrap() + } + + async fn retire(registry: &PoolRegistry, attachment: PoolAttachment, mode: PoolRetirementMode) { + let PoolAttachment { + pool_id, + layout_generation, + handle, + mut faults, + .. + } = attachment; + assert!(registry.withdraw(pool_id, layout_generation, mode).await); + let result = match mode { + PoolRetirementMode::Graceful => { + drain_faults_while(pool_id, &mut faults, handle.shutdown()).await + } + PoolRetirementMode::Fenced => { + drain_faults_while(pool_id, &mut faults, handle.fence()).await + } + }; + result.unwrap(); + assert!(registry.remove(pool_id, layout_generation).await); + } + + #[tokio::test] + async fn one_model_binds_to_independent_pools() { + let registry = PoolRegistry::new(relay_identity(), config()); + let first = registry + .attach(request(pool(1), "fast.router.generate", "llama")) + .await + .unwrap(); + let second = registry + .attach(request(pool(2), "slow.router.generate", "llama")) + .await + .unwrap(); + + assert_eq!(registry.pool_count().await, 2); + let catalog = registry.catalog(); + assert_eq!(catalog.drt_instance_id(), 11); + assert_eq!(catalog.relay_incarnation(), 7); + assert_eq!(catalog.pools().len(), 2); + assert_eq!( + descriptor(&catalog, pool(1)).serving_endpoint(), + &EndpointId::from("fast.router.generate") + ); + assert!( + catalog + .pools() + .iter() + .all(|descriptor| { descriptor.registrations()[0].model().as_str() == "llama" }) + ); + + registry.detach(first).await.unwrap(); + registry.detach(second).await.unwrap(); + } + + #[tokio::test] + async fn burst_attach_defers_catalog_materialization_until_observed() { + let registry = PoolRegistry::new(relay_identity(), config()); + let catalog_rx = registry.watch_catalog(); + let mut attachments = Vec::new(); + + for seed in 1..=32 { + attachments.push( + registry + .attach(request( + pool(seed), + &format!("pool-{seed}.router.generate"), + &format!("model-{seed}"), + )) + .await + .unwrap(), + ); + } + + assert_eq!(catalog_rx.borrow().revision(), 32); + assert!(!catalog_rx.borrow().is_materialized()); + + let catalog = registry.catalog(); + let pool_ids: Vec<_> = catalog + .pools() + .iter() + .map(DcPoolDescriptor::pool_id) + .collect(); + assert!(catalog.is_materialized()); + assert!(pool_ids.windows(2).all(|pair| pair[0] < pair[1])); + let serialized = serde_json::to_value(&catalog).unwrap(); + assert_eq!(serialized["drt_instance_id"], 11); + assert_eq!(serialized["relay_incarnation"], 7); + assert!(serialized.get("process_incarnation").is_none()); + assert_eq!(serialized["revision"], 32); + assert_eq!(serialized["pools"].as_array().unwrap().len(), 32); + + registry.shutdown().await; + assert_eq!(catalog.pools().len(), attachments.len()); + assert!(registry.catalog().pools().is_empty()); + } + + #[tokio::test] + async fn registration_updates_remain_pool_local() { + let registry = PoolRegistry::new(relay_identity(), config()); + let mut first = registry + .attach(request(pool(1), "fast.router.generate", "llama")) + .await + .unwrap(); + let mut second = registry + .attach(request(pool(2), "slow.router.generate", "mistral")) + .await + .unwrap(); + + registry + .replace_registrations( + &mut first, + vec![CanonicalModelRegistration::new( + CanonicalModelId::new("llama").unwrap(), + vec![ModelAlias::new("mistral-alias").unwrap()], + )], + ) + .await + .unwrap(); + let catalog = registry.catalog(); + assert!(catalog.pools().iter().all(|descriptor| { + descriptor.registrations()[0].aliases() == [ModelAlias::new("mistral-alias").unwrap()] + })); + + registry + .replace_registrations( + &mut second, + vec![CanonicalModelRegistration::new( + CanonicalModelId::new("mistral").unwrap(), + vec![ModelAlias::new("llama-alias").unwrap()], + )], + ) + .await + .unwrap(); + let catalog = registry.catalog(); + assert_eq!( + descriptor(&catalog, pool(1)).registrations()[0].aliases(), + [ModelAlias::new("mistral-alias").unwrap()] + ); + assert_eq!( + descriptor(&catalog, pool(2)).registrations()[0].aliases(), + [ModelAlias::new("llama-alias").unwrap()] + ); + + registry.detach(first).await.unwrap(); + registry.detach(second).await.unwrap(); + } + + #[tokio::test] + async fn one_pool_cannot_be_owned_by_two_endpoints() { + let registry = PoolRegistry::new(relay_identity(), config()); + let pool_id = pool(1); + let attachment = registry + .attach(request(pool_id, "first.router.generate", "llama")) + .await + .unwrap(); + + let error = registry + .attach(request(pool_id, "second.router.generate", "llama")) + .await + .err() + .unwrap(); + assert!(error.to_string().contains("already owned")); + + registry.detach(attachment).await.unwrap(); + } + + #[tokio::test] + async fn cancelled_allocation_rolls_back_and_allows_reattach() { + let registry = Arc::new(PoolRegistry::new(relay_identity(), config())); + let pool_id = pool(1); + let GatedBuilder { + builder, + started: started_rx, + release: release_tx, + finished: finished_rx, + } = gated_builder(); + let task_registry = registry.clone(); + let attach = tokio::spawn(async move { + task_registry + .attach_with_builder(request(pool_id, "first.router.generate", "llama"), builder) + .await + }); + + started_rx.await.unwrap(); + attach.abort(); + assert!(matches!(attach.await, Err(error) if error.is_cancelled())); + assert!(registry.state.lock().reservations.is_empty()); + + release_tx.send(()).unwrap(); + finished_rx.await.unwrap(); + let attachment = registry + .attach(request(pool_id, "second.router.generate", "llama")) + .await + .unwrap(); + registry.detach(attachment).await.unwrap(); + } + + #[tokio::test] + async fn shutdown_during_allocation_never_publishes_the_pool() { + let registry = Arc::new(PoolRegistry::new(relay_identity(), config())); + let GatedBuilder { + builder, + started: started_rx, + release: release_tx, + finished: finished_rx, + } = gated_builder(); + let task_registry = registry.clone(); + let attach = tokio::spawn(async move { + task_registry + .attach_with_builder(request(pool(1), "first.router.generate", "llama"), builder) + .await + }); + + started_rx.await.unwrap(); + registry.shutdown().await; + assert!(registry.state.lock().reservations.is_empty()); + assert!(registry.catalog().pools().is_empty()); + + release_tx.send(()).unwrap(); + finished_rx.await.unwrap(); + let Err(error) = attach.await.unwrap() else { + panic!("pool attached after registry shutdown"); + }; + assert!(error.to_string().contains("retired before commit")); + assert_eq!(registry.pool_count().await, 0); + assert!(registry.catalog().pools().is_empty()); + } + + #[tokio::test(flavor = "current_thread")] + async fn ckf_allocation_does_not_block_the_async_executor() { + let registry = Arc::new(PoolRegistry::new(relay_identity(), config())); + let GatedBuilder { + builder, + started: started_rx, + release: release_tx, + finished: finished_rx, + } = gated_builder(); + let task_registry = registry.clone(); + let attach = tokio::spawn(async move { + task_registry + .attach_with_builder(request(pool(1), "first.router.generate", "llama"), builder) + .await + }); + + tokio::time::timeout(Duration::from_millis(500), started_rx) + .await + .expect("blocking CKF allocation did not start") + .unwrap(); + tokio::time::timeout( + Duration::from_millis(500), + tokio::time::sleep(Duration::from_millis(1)), + ) + .await + .expect("CKF allocation blocked the async executor"); + + release_tx.send(()).unwrap(); + finished_rx.await.unwrap(); + let attachment = attach.await.unwrap().unwrap(); + registry.detach(attachment).await.unwrap(); + } + + #[tokio::test] + async fn failed_actor_build_rolls_back_the_pool_reservation() { + let registry = PoolRegistry::new( + relay_identity(), + PoolActorConfig { + expected_unique_blocks: 0, + publication_threshold: 1, + publication_delay: Duration::from_millis(1), + }, + ); + let pool_id = pool(1); + + let first_error = registry + .attach(request(pool_id, "first.router.generate", "llama")) + .await + .err() + .unwrap(); + let second_error = registry + .attach(request(pool_id, "second.router.generate", "llama")) + .await + .err() + .unwrap(); + + assert!(first_error.to_string().contains("greater than zero")); + assert!(second_error.to_string().contains("greater than zero")); + assert_eq!(registry.pool_count().await, 0); + assert_eq!(registry.catalog().revision(), 0); + } + + #[tokio::test] + async fn reattaching_a_pool_allocates_a_new_layout_generation() { + let registry = PoolRegistry::new(relay_identity(), config()); + let pool_id = pool(1); + let first = registry + .attach(request(pool_id, "fast.router.generate", "llama")) + .await + .unwrap(); + let first_generation = first.layout_generation; + registry.detach(first).await.unwrap(); + + let replacement = registry + .attach(request(pool_id, "fast.router.generate", "llama")) + .await + .unwrap(); + assert!(replacement.layout_generation > first_generation); + + registry.detach(replacement).await.unwrap(); + } + + #[tokio::test] + async fn relay_incarnation_fences_an_identical_pool_layout() { + let first_registry = PoolRegistry::new(DcRelayIdentity::new(11, 7), config()); + let second_registry = PoolRegistry::new(DcRelayIdentity::new(11, 8), config()); + let pool_id = pool(1); + let first = first_registry + .attach(request(pool_id, "fast.router.generate", "llama")) + .await + .unwrap(); + let second = second_registry + .attach(request(pool_id, "fast.router.generate", "llama")) + .await + .unwrap(); + + assert_ne!(first.handle.identity(), second.handle.identity()); + assert_eq!(first.layout_generation, second.layout_generation); + + first_registry.detach(first).await.unwrap(); + second_registry.detach(second).await.unwrap(); + } + + #[tokio::test] + async fn withdraw_removes_catalog_before_actor_retirement() { + let registry = PoolRegistry::new(relay_identity(), config()); + let attachment = registry + .attach(request(pool(1), "fast.router.generate", "llama")) + .await + .unwrap(); + + assert!( + registry + .withdraw( + attachment.pool_id, + attachment.layout_generation, + PoolRetirementMode::Graceful, + ) + .await + ); + assert!(registry.catalog().pools().is_empty()); + attachment.handle.state_stats().await.unwrap(); + + registry.detach(attachment).await.unwrap(); + } + + #[tokio::test] + async fn adapter_registration_changes_without_replacing_pool_generation() { + let registry = PoolRegistry::new(relay_identity(), config()); + let mut attachment = registry + .attach(request(pool(1), "fast.router.generate", "llama")) + .await + .unwrap(); + let generation = attachment.layout_generation; + let base = CanonicalModelId::new("llama").unwrap(); + let adapter = CanonicalModelId::new("tenant-a").unwrap(); + registry + .replace_registrations( + &mut attachment, + vec![ + CanonicalModelRegistration::new(base.clone(), Vec::new()), + CanonicalModelRegistration::with_target( + adapter.clone(), + ModelTarget::Lora { + base_model: base, + adapter: adapter.clone(), + }, + Vec::new(), + ), + ], + ) + .await + .unwrap(); + + assert_eq!(attachment.layout_generation, generation); + let catalog = registry.catalog(); + assert!( + descriptor(&catalog, pool(1)) + .registrations() + .iter() + .any(|registration| registration.model() == &adapter) + ); + registry.detach(attachment).await.unwrap(); + } + + #[tokio::test] + async fn one_lora_target_binds_to_independent_pools() { + let registry = PoolRegistry::new(relay_identity(), config()); + let base = CanonicalModelId::new("llama").unwrap(); + let adapter = CanonicalModelId::new("tenant-a").unwrap(); + let registrations = || { + vec![ + CanonicalModelRegistration::new(base.clone(), Vec::new()), + CanonicalModelRegistration::with_target( + adapter.clone(), + ModelTarget::Lora { + base_model: base.clone(), + adapter: adapter.clone(), + }, + Vec::new(), + ), + ] + }; + let first = registry + .attach(PoolAttachRequest { + pool_id: pool(1), + endpoint: EndpointId::from("fast.router.generate"), + registrations: registrations(), + }) + .await + .unwrap(); + let second = registry + .attach(PoolAttachRequest { + pool_id: pool(2), + endpoint: EndpointId::from("slow.router.generate"), + registrations: registrations(), + }) + .await + .unwrap(); + + let catalog = registry.catalog(); + assert_eq!(catalog.pools().len(), 2); + assert!(catalog.pools().iter().all(|descriptor| { + descriptor.registrations().iter().any(|registration| { + registration.target() + == &ModelTarget::Lora { + base_model: base.clone(), + adapter: adapter.clone(), + } + }) + })); + + registry.detach(first).await.unwrap(); + registry.detach(second).await.unwrap(); + } + + #[tokio::test] + async fn fencing_withdraws_pool_from_catalog() { + let registry = PoolRegistry::new(relay_identity(), config()); + let attachment = registry + .attach(request(pool(1), "fast.router.generate", "llama")) + .await + .unwrap(); + retire(®istry, attachment, PoolRetirementMode::Fenced).await; + assert!(registry.watch_catalog().borrow().pools().is_empty()); + } + + #[tokio::test] + async fn fencing_withdraws_only_the_target_pool_descriptor() { + let registry = PoolRegistry::new(relay_identity(), config()); + let with_alias = registry + .attach(request(pool(1), "fast.router.generate", "llama")) + .await + .unwrap(); + let without_alias = registry + .attach(PoolAttachRequest { + pool_id: pool(2), + endpoint: EndpointId::from("slow.router.generate"), + registrations: vec![CanonicalModelRegistration::new( + CanonicalModelId::new("llama").unwrap(), + Vec::new(), + )], + }) + .await + .unwrap(); + + let catalog = registry.catalog(); + assert_eq!( + descriptor(&catalog, pool(1)).registrations()[0] + .aliases() + .len(), + 1 + ); + assert!( + descriptor(&catalog, pool(2)).registrations()[0] + .aliases() + .is_empty() + ); + retire(®istry, with_alias, PoolRetirementMode::Fenced).await; + let catalog = registry.catalog(); + assert_eq!(catalog.pools().len(), 1); + assert_eq!(catalog.pools()[0].pool_id(), pool(2)); + assert!(catalog.pools()[0].registrations()[0].aliases().is_empty()); + + registry.detach(without_alias).await.unwrap(); + } +} diff --git a/lib/llm/src/kv_dc_relay/resolution.rs b/lib/llm/src/kv_dc_relay/resolution.rs index 21f4cbdbcae5..39b159f4fb9c 100644 --- a/lib/llm/src/kv_dc_relay/resolution.rs +++ b/lib/llm/src/kv_dc_relay/resolution.rs @@ -46,6 +46,7 @@ pub(crate) struct EndpointLocator { endpoint_id: EndpointId, } +#[allow(dead_code)] impl EndpointLocator { pub(crate) fn new(dc_id: DcId, endpoint_id: EndpointId) -> Self { Self { dc_id, endpoint_id } @@ -66,6 +67,7 @@ pub(crate) struct PoolBinding { kv_state_endpoint: Option, } +#[allow(dead_code)] impl PoolBinding { pub(crate) fn new( pool_id: PoolId,