Skip to main content

dynamo_runtime/discovery/
kube.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4mod crd;
5mod daemon;
6mod utils;
7
8pub use crd::{DynamoWorkerMetadata, DynamoWorkerMetadataSpec};
9// hash_pod_name/hash_container_name are used by C bindings and the Rust EPP
10// for pod- and container-level worker ID mapping.
11pub use utils::{hash_container_name, hash_pod_name};
12
13use crd::{apply_cr, build_cr};
14use daemon::DiscoveryDaemon;
15use utils::{KubeDiscoveryMode, PodInfo};
16
17use crate::CancellationToken;
18use crate::discovery::{
19    Discovery, DiscoveryEvent, DiscoveryInstance, DiscoveryInstanceId, DiscoveryMetadata,
20    DiscoveryQuery, DiscoverySpec, DiscoveryStream, MAX_JSON_SAFE_PUBLISHER_ID, MetadataSnapshot,
21    ModelCardInstanceId, reconcile_discovery_snapshot,
22};
23use anyhow::Result;
24use async_trait::async_trait;
25use kube::{Api, Client as KubeClient, api::DeleteParams};
26use std::collections::{HashMap, HashSet};
27use std::future::Future;
28use std::sync::Arc;
29use tokio::sync::RwLock;
30
31fn validate_kubernetes_publisher_id(publisher_id: u64) -> Result<()> {
32    if publisher_id > MAX_JSON_SAFE_PUBLISHER_ID {
33        anyhow::bail!(
34            "Kubernetes discovery publisher ID {publisher_id} exceeds the JSON-safe maximum \
35             {MAX_JSON_SAFE_PUBLISHER_ID}"
36        );
37    }
38
39    Ok(())
40}
41
42async fn update_model_taints_and_persist<F, Fut>(
43    metadata: &Arc<RwLock<DiscoveryMetadata>>,
44    id: ModelCardInstanceId,
45    taints: HashSet<String>,
46    persist: F,
47) -> Result<bool>
48where
49    F: FnOnce(DiscoveryMetadata) -> Fut + Send + 'static,
50    Fut: Future<Output = Result<DiscoveryMetadata>> + Send + 'static,
51{
52    let metadata = Arc::clone(metadata);
53    // Once started, persistence and the matching local commit must outlive request cancellation.
54    // Dropping the JoinHandle detaches this task instead of cancelling the remote-commit/local-
55    // state critical section.
56    tokio::spawn(async move {
57        let mut metadata = metadata.write().await;
58        let mut candidate = metadata.clone();
59        let changed = candidate.update_model_taints(&id, taints)?;
60
61        // Persist even a local no-op. This repairs an authoritative CR that may differ after an
62        // earlier commit/ack ambiguity instead of trusting potentially stale local metadata.
63        let persisted = persist(candidate).await?;
64        *metadata = persisted;
65        Ok(changed)
66    })
67    .await
68    .map_err(|error| anyhow::anyhow!("model taint persistence task failed: {error}"))?
69}
70
71/// Kubernetes-based discovery client
72#[derive(Clone)]
73pub struct KubeDiscoveryClient {
74    instance_id: u64,
75    metadata: Arc<RwLock<DiscoveryMetadata>>,
76    metadata_watch: tokio::sync::watch::Receiver<Arc<MetadataSnapshot>>,
77    kube_client: KubeClient,
78    pod_info: PodInfo,
79}
80
81impl KubeDiscoveryClient {
82    /// Create a new Kubernetes discovery client
83    ///
84    /// # Arguments
85    /// * `metadata` - Shared metadata store (also used by system server)
86    /// * `cancel_token` - Cancellation token for shutdown
87    pub async fn new(
88        metadata: Arc<RwLock<DiscoveryMetadata>>,
89        cancel_token: CancellationToken,
90    ) -> Result<Self> {
91        let pod_info = PodInfo::from_env()?;
92        let instance_id = pod_info.target.instance_id();
93        let cr_name = pod_info.target.cr_name();
94
95        tracing::info!(
96            "Initializing KubeDiscoveryClient: mode={:?}, target={:?}, cr_name={}, instance_id={:x}, namespace={}, pod_uid={}",
97            pod_info.mode,
98            pod_info.target,
99            cr_name,
100            instance_id,
101            pod_info.pod_namespace,
102            pod_info.pod_uid
103        );
104
105        let kube_client = KubeClient::try_default()
106            .await
107            .map_err(|e| anyhow::anyhow!("Failed to create Kubernetes client: {}", e))?;
108
109        // In container mode, delete any stale CR from a previous incarnation of this container.
110        // In failover pods, the pod stays alive when a container crashes and restarts,
111        // so the old CR persists. Deleting it ensures the daemon doesn't see stale data.
112        // In pod mode this is unnecessary — pod restart creates a new pod (and new CR name).
113        if pod_info.mode == KubeDiscoveryMode::Container {
114            let cr_api: Api<DynamoWorkerMetadata> =
115                Api::namespaced(kube_client.clone(), &pod_info.pod_namespace);
116            match cr_api.delete(&cr_name, &DeleteParams::default()).await {
117                Ok(_) => tracing::info!("Deleted stale CR: {}", cr_name),
118                Err(kube::Error::Api(err_resp)) if err_resp.code == 404 => {
119                    tracing::debug!("No stale CR to delete: {}", cr_name);
120                }
121                Err(e) => {
122                    panic!(
123                        "Failed to clear stale CR '{}': {} — cannot start with stale discovery state",
124                        cr_name, e
125                    );
126                }
127            }
128        }
129
130        // Create watch channel with initial empty snapshot
131        let (watch_tx, watch_rx) = tokio::sync::watch::channel(Arc::new(MetadataSnapshot::empty()));
132
133        // Create and spawn daemon
134        let daemon = DiscoveryDaemon::new(kube_client.clone(), pod_info.clone(), cancel_token)?;
135
136        tokio::spawn(async move {
137            if let Err(e) = daemon.run(watch_tx).await {
138                tracing::error!("Discovery daemon failed: {e}");
139            }
140        });
141
142        tracing::info!("Discovery daemon started");
143
144        Ok(Self {
145            instance_id,
146            metadata,
147            metadata_watch: watch_rx,
148            kube_client,
149            pod_info,
150        })
151    }
152}
153
154#[async_trait]
155impl Discovery for KubeDiscoveryClient {
156    fn instance_id(&self) -> u64 {
157        self.instance_id
158    }
159
160    async fn register_internal(&self, spec: DiscoverySpec) -> Result<DiscoveryInstance> {
161        match &spec {
162            DiscoverySpec::EventChannel { publisher_id, .. }
163            | DiscoverySpec::EventSource { publisher_id, .. } => {
164                validate_kubernetes_publisher_id(*publisher_id)?;
165            }
166            _ => {}
167        }
168        let instance = spec.into_instance(self.instance_id());
169        let instance_id = instance.instance_id();
170
171        tracing::debug!(
172            "Registering discovery instance: {:?}, instance_id={:x}",
173            instance,
174            instance_id
175        );
176
177        // Write to local metadata and persist to CR
178        // IMPORTANT: Hold the write lock across the CR write to prevent race conditions
179        let mut metadata = self.metadata.write().await;
180
181        // Clone state for rollback in case CR persistence fails
182        let original_state = metadata.clone();
183
184        let registered_instance = match &instance {
185            DiscoveryInstance::Endpoint(inst) => {
186                tracing::info!(
187                    "Registering endpoint: namespace={}, component={}, endpoint={}, instance_id={:x}",
188                    inst.namespace,
189                    inst.component,
190                    inst.endpoint,
191                    instance_id
192                );
193                metadata.register_endpoint(instance.clone())?;
194                instance.clone()
195            }
196            DiscoveryInstance::Model {
197                namespace,
198                component,
199                endpoint,
200                ..
201            } => {
202                tracing::info!(
203                    "Registering model card: namespace={}, component={}, endpoint={}, instance_id={:x}",
204                    namespace,
205                    component,
206                    endpoint,
207                    instance_id
208                );
209                metadata.register_model_card(instance.clone())?
210            }
211            DiscoveryInstance::EventChannel { scope, topic, .. } => {
212                tracing::info!(
213                    "Registering event channel: scope={:?}, topic={}, instance_id={:x}",
214                    scope,
215                    topic,
216                    instance_id
217                );
218                metadata.register_event_channel(instance.clone())?;
219                instance.clone()
220            }
221            DiscoveryInstance::EventSource { scope, topic, .. } => {
222                tracing::info!(
223                    "Registering event source: scope={:?}, topic={}, publisher_id={:x}",
224                    scope,
225                    topic,
226                    instance_id
227                );
228                metadata.register_event_source(instance.clone())?;
229                instance.clone()
230            }
231        };
232
233        // Build and apply the CR with the updated metadata
234        // This persists the metadata to Kubernetes for other pods to discover
235        let cr_name = self.pod_info.target.cr_name();
236        let cr = build_cr(
237            &cr_name,
238            &self.pod_info.pod_name,
239            &self.pod_info.pod_uid,
240            &metadata,
241        )?;
242
243        if let Err(e) = apply_cr(&self.kube_client, &self.pod_info.pod_namespace, &cr).await {
244            // Rollback local state on CR persistence failure
245            tracing::warn!(
246                "Failed to persist metadata to CR, rolling back local state: {}",
247                e
248            );
249            *metadata = original_state;
250            return Err(e);
251        }
252
253        tracing::debug!("Persisted metadata to DynamoWorkerMetadata CR");
254
255        Ok(registered_instance)
256    }
257
258    async fn update_model_taints_internal(
259        &self,
260        id: ModelCardInstanceId,
261        taints: HashSet<String>,
262    ) -> Result<()> {
263        let kube_client = self.kube_client.clone();
264        let pod_namespace = self.pod_info.pod_namespace.clone();
265        let cr_name = self.pod_info.target.cr_name();
266        let pod_name = self.pod_info.pod_name.clone();
267        let pod_uid = self.pod_info.pod_uid.clone();
268        let changed = update_model_taints_and_persist(
269            &self.metadata,
270            id,
271            taints,
272            move |candidate| async move {
273                let cr = build_cr(&cr_name, &pod_name, &pod_uid, &candidate)?;
274                apply_cr(&kube_client, &pod_namespace, &cr).await?;
275                Ok(candidate)
276            },
277        )
278        .await?;
279        if !changed {
280            return Ok(());
281        }
282
283        tracing::debug!("Persisted model taint update to DynamoWorkerMetadata CR");
284        Ok(())
285    }
286
287    async fn unregister(&self, instance: DiscoveryInstance) -> Result<()> {
288        let instance_id = instance.instance_id();
289
290        // Write to local metadata and persist to CR
291        // IMPORTANT: Hold the write lock across the CR write to prevent race conditions
292        let mut metadata = self.metadata.write().await;
293
294        // Clone state for rollback in case CR persistence fails
295        let original_state = metadata.clone();
296
297        match &instance {
298            DiscoveryInstance::Endpoint(inst) => {
299                tracing::info!(
300                    "Unregistering endpoint: namespace={}, component={}, endpoint={}, instance_id={:x}",
301                    inst.namespace,
302                    inst.component,
303                    inst.endpoint,
304                    instance_id
305                );
306                metadata.unregister_endpoint(&instance)?;
307            }
308            DiscoveryInstance::Model {
309                namespace,
310                component,
311                endpoint,
312                ..
313            } => {
314                tracing::info!(
315                    "Unregistering model card: namespace={}, component={}, endpoint={}, instance_id={:x}",
316                    namespace,
317                    component,
318                    endpoint,
319                    instance_id
320                );
321                metadata.unregister_model_card(&instance)?;
322            }
323            DiscoveryInstance::EventChannel { scope, topic, .. } => {
324                tracing::info!(
325                    "Unregistering event channel: scope={:?}, topic={}, instance_id={:x}",
326                    scope,
327                    topic,
328                    instance_id
329                );
330                metadata.unregister_event_channel(&instance)?;
331            }
332            DiscoveryInstance::EventSource { scope, topic, .. } => {
333                tracing::info!(
334                    "Unregistering event source: scope={:?}, topic={}, publisher_id={:x}",
335                    scope,
336                    topic,
337                    instance_id
338                );
339                metadata.unregister_event_source(&instance)?;
340            }
341        }
342
343        // Build and apply the CR with the updated metadata
344        // This persists the removal to Kubernetes for other pods to see
345        let cr_name = self.pod_info.target.cr_name();
346        let cr = build_cr(
347            &cr_name,
348            &self.pod_info.pod_name,
349            &self.pod_info.pod_uid,
350            &metadata,
351        )?;
352
353        if let Err(e) = apply_cr(&self.kube_client, &self.pod_info.pod_namespace, &cr).await {
354            // Rollback local state on CR persistence failure
355            tracing::warn!(
356                "Failed to persist metadata removal to CR, rolling back local state: {}",
357                e
358            );
359            *metadata = original_state;
360            return Err(e);
361        }
362
363        tracing::debug!("Persisted metadata removal to DynamoWorkerMetadata CR");
364
365        Ok(())
366    }
367
368    async fn list(&self, query: DiscoveryQuery) -> Result<Vec<DiscoveryInstance>> {
369        tracing::debug!("KubeDiscoveryClient::list called with query={:?}", query);
370
371        // Get current snapshot (may be empty if daemon hasn't fetched yet)
372        let snapshot = self.metadata_watch.borrow().clone();
373
374        tracing::debug!(
375            "List using snapshot seq={} with {} instances",
376            snapshot.sequence,
377            snapshot.instances.len()
378        );
379
380        // Filter snapshot by query
381        let instances = snapshot.filter(&query);
382
383        tracing::info!(
384            "KubeDiscoveryClient::list returning {} instances for query={:?}",
385            instances.len(),
386            query
387        );
388
389        Ok(instances)
390    }
391
392    async fn list_and_watch(
393        &self,
394        query: DiscoveryQuery,
395        cancel_token: Option<CancellationToken>,
396    ) -> Result<DiscoveryStream> {
397        use tokio::sync::mpsc;
398
399        tracing::info!(
400            "KubeDiscoveryClient::list_and_watch started for query={:?}",
401            query
402        );
403
404        // Clone the watch receiver
405        let mut watch_rx = self.metadata_watch.clone();
406
407        // Create output stream
408        let (event_tx, event_rx) = mpsc::unbounded_channel();
409
410        // Generate unique stream identifier for tracing
411        let stream_id = uuid::Uuid::new_v4();
412
413        // Spawn task to process snapshots
414        tokio::spawn(async move {
415            // Initialize from current snapshot state
416            // This is critical: watch_rx.changed() only fires on FUTURE changes,
417            // so we must capture the current state first to detect removals correctly
418            let initial_snapshot = watch_rx.borrow_and_update().clone();
419
420            // Build initial map: DiscoveryInstanceId -> DiscoveryInstance
421            let initial: HashMap<DiscoveryInstanceId, DiscoveryInstance> = initial_snapshot
422                .instances
423                .values()
424                .flat_map(|metadata| metadata.filter(&query))
425                .map(|instance| (instance.id(), instance))
426                .collect();
427
428            tracing::debug!(
429                stream_id = %stream_id,
430                initial_count = initial.len(),
431                "Watch started for query={:?}",
432                query
433            );
434
435            // Emit initial Added events (the "list" part of list_and_watch)
436            for instance in initial.values() {
437                tracing::info!(
438                    stream_id = %stream_id,
439                    instance_id = format!("{:x}", instance.instance_id()),
440                    "Emitting initial Added event"
441                );
442                if event_tx
443                    .send(Ok(DiscoveryEvent::Added(instance.clone())))
444                    .is_err()
445                {
446                    tracing::debug!(
447                        stream_id = %stream_id,
448                        "Watch receiver dropped during initial sync"
449                    );
450                    return;
451                }
452            }
453
454            // Track complete values so same-ID model taint updates are observable.
455            let mut known = initial;
456
457            loop {
458                tracing::trace!(
459                    stream_id = %stream_id,
460                    known_count = known.len(),
461                    "Watch loop waiting for changes"
462                );
463
464                // Wait for next snapshot or cancellation
465                let watch_result = if let Some(ref token) = cancel_token {
466                    tokio::select! {
467                        result = watch_rx.changed() => result,
468                        _ = token.cancelled() => {
469                            tracing::info!(
470                                stream_id = %stream_id,
471                                "Watch cancelled via cancel token"
472                            );
473                            break;
474                        }
475                    }
476                } else {
477                    watch_rx.changed().await
478                };
479
480                match watch_result {
481                    Ok(()) => {
482                        // Get latest snapshot
483                        let snapshot = watch_rx.borrow_and_update().clone();
484
485                        // Build current map: DiscoveryInstanceId -> DiscoveryInstance
486                        let current: HashMap<DiscoveryInstanceId, DiscoveryInstance> = snapshot
487                            .instances
488                            .values()
489                            .flat_map(|metadata| metadata.filter(&query))
490                            .map(|instance| (instance.id(), instance))
491                            .collect();
492
493                        tracing::debug!(
494                            stream_id = %stream_id,
495                            seq = snapshot.sequence,
496                            current_count = current.len(),
497                            known_count = known.len(),
498                            "Watch received snapshot update"
499                        );
500
501                        let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
502
503                        // Log diff results (even if empty, for debugging)
504                        if events.is_empty() {
505                            tracing::debug!(
506                                stream_id = %stream_id,
507                                seq = snapshot.sequence,
508                                "Watch snapshot received but no diff detected"
509                            );
510                        } else {
511                            tracing::debug!(
512                                stream_id = %stream_id,
513                                seq = snapshot.sequence,
514                                emitted_events = events.len(),
515                                total = reconciled.len(),
516                                "Watch detected changes"
517                            );
518                        }
519
520                        for event in events {
521                            let (event_kind, instance_id) = match &event {
522                                DiscoveryEvent::Added(instance) => ("added", instance.id()),
523                                DiscoveryEvent::ModelTaintsUpdated(update) => (
524                                    "model_taints_updated",
525                                    DiscoveryInstanceId::Model(update.id.clone()),
526                                ),
527                                DiscoveryEvent::Removed(id) => ("removed", id.clone()),
528                            };
529                            tracing::info!(
530                                stream_id = %stream_id,
531                                event_kind,
532                                ?instance_id,
533                                "Emitting discovery event"
534                            );
535                            tracing::debug!(
536                                stream_id = %stream_id,
537                                ?event,
538                                "Discovery event detail"
539                            );
540                            if event_tx.send(Ok(event)).is_err() {
541                                tracing::debug!(stream_id = %stream_id, "Watch receiver dropped");
542                                return;
543                            }
544                        }
545
546                        known = reconciled;
547                    }
548                    Err(_) => {
549                        tracing::info!(
550                            stream_id = %stream_id,
551                            "Watch channel closed (daemon stopped)"
552                        );
553                        break;
554                    }
555                }
556            }
557        });
558
559        // Convert receiver to stream
560        let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(event_rx);
561        Ok(Box::pin(stream))
562    }
563}
564
565#[cfg(test)]
566mod tests {
567    use super::*;
568    use crate::component::TransportType;
569    use crate::discovery::{EventScope, EventTransport, ModelTaintsUpdate};
570
571    fn endpoint_instance(instance_id: u64, transport: &str) -> DiscoveryInstance {
572        DiscoveryInstance::Endpoint(crate::component::Instance {
573            namespace: "ns".to_string(),
574            component: "component".to_string(),
575            endpoint: "endpoint".to_string(),
576            instance_id,
577            transport: TransportType::Tcp(transport.to_string()),
578            device_type: None,
579            request_plane_codec: None,
580        })
581    }
582
583    fn model_with_taint(taint: &str) -> DiscoveryInstance {
584        DiscoveryInstance::Model {
585            namespace: "ns".to_string(),
586            component: "worker".to_string(),
587            endpoint: "generate".to_string(),
588            instance_id: 7,
589            card_json: serde_json::json!({
590                "runtime_config": {"taints": [taint]}
591            }),
592            model_suffix: None,
593        }
594    }
595
596    #[test]
597    fn publisher_ids_must_fit_kubernetes_json_safe_range() {
598        assert!(validate_kubernetes_publisher_id(MAX_JSON_SAFE_PUBLISHER_ID).is_ok());
599        assert!(validate_kubernetes_publisher_id(MAX_JSON_SAFE_PUBLISHER_ID + 1).is_err());
600        assert!(validate_kubernetes_publisher_id(u64::MAX).is_err());
601    }
602
603    #[test]
604    fn snapshot_diff_emits_updated_instance_when_transport_changes() {
605        let original = endpoint_instance(1, "127.0.0.1:8000");
606        let updated = endpoint_instance(1, "127.0.0.1:9000");
607        let known = HashMap::from([(original.id(), original)]);
608        let current = HashMap::from([(updated.id(), updated.clone())]);
609
610        let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
611
612        assert_eq!(events, vec![DiscoveryEvent::Added(updated.clone())]);
613        assert_eq!(reconciled.get(&updated.id()), Some(&updated));
614    }
615
616    #[test]
617    fn snapshot_diff_ignores_same_id_event_channel_changes() {
618        let event_channel = |endpoint: &str| DiscoveryInstance::EventChannel {
619            scope: EventScope::Namespace {
620                name: "ns".to_string(),
621            },
622            topic: "topic".to_string(),
623            instance_id: 1,
624            transport: EventTransport::zmq(endpoint),
625        };
626        let original = event_channel("tcp://127.0.0.1:8000");
627        let updated = event_channel("tcp://127.0.0.1:9000");
628        let known = HashMap::from([(original.id(), original.clone())]);
629        let current = HashMap::from([(updated.id(), updated)]);
630
631        let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
632
633        assert!(events.is_empty());
634        assert_eq!(reconciled.get(&original.id()), Some(&original));
635    }
636
637    #[test]
638    fn snapshot_diff_emits_added_and_removed_instances() {
639        let removed_instance = endpoint_instance(1, "127.0.0.1:8000");
640        let added_instance = endpoint_instance(2, "127.0.0.1:9000");
641        let removed_id = removed_instance.id();
642        let added_id = added_instance.id();
643        let known = HashMap::from([(removed_id.clone(), removed_instance)]);
644        let current = HashMap::from([(added_id.clone(), added_instance.clone())]);
645
646        let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
647
648        assert_eq!(events.len(), 2);
649        assert!(events.contains(&DiscoveryEvent::Removed(removed_id.clone())));
650        assert!(events.contains(&DiscoveryEvent::Added(added_instance.clone())));
651        assert!(!reconciled.contains_key(&removed_id));
652        assert_eq!(reconciled.get(&added_id), Some(&added_instance));
653    }
654    #[test]
655    fn changed_model_taints_emit_scoped_event() {
656        let old = model_with_taint("old");
657        let updated = model_with_taint("updated");
658        let known = HashMap::from([(old.id(), old)]);
659        let current = HashMap::from([(updated.id(), updated.clone())]);
660
661        let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
662
663        let DiscoveryInstanceId::Model(id) = updated.id() else {
664            unreachable!()
665        };
666        assert_eq!(
667            events,
668            vec![DiscoveryEvent::ModelTaintsUpdated(ModelTaintsUpdate {
669                id,
670                taints: vec!["updated".to_string()],
671            })]
672        );
673        assert_eq!(reconciled.get(&updated.id()), Some(&updated));
674    }
675
676    #[tokio::test]
677    async fn model_taint_persistence_completes_after_caller_cancellation() {
678        let model = model_with_taint("old");
679        let DiscoveryInstanceId::Model(id) = model.id() else {
680            unreachable!()
681        };
682        let mut initial = DiscoveryMetadata::new();
683        initial.register_model_card(model).unwrap();
684        let metadata = Arc::new(RwLock::new(initial));
685        let task_metadata = metadata.clone();
686        let remote = Arc::new(RwLock::new(DiscoveryMetadata::new()));
687        let task_remote = remote.clone();
688        let (remote_committed_tx, remote_committed_rx) = tokio::sync::oneshot::channel();
689        let (ack_tx, ack_rx) = tokio::sync::oneshot::channel();
690
691        let task = tokio::spawn(async move {
692            update_model_taints_and_persist(
693                &task_metadata,
694                id,
695                HashSet::from(["new".to_string()]),
696                move |candidate| async move {
697                    *task_remote.write().await = candidate.clone();
698                    remote_committed_tx.send(()).unwrap();
699                    ack_rx.await.unwrap();
700                    Ok(candidate)
701                },
702            )
703            .await
704        });
705
706        remote_committed_rx.await.unwrap();
707        task.abort();
708        assert!(task.await.unwrap_err().is_cancelled());
709        ack_tx.send(()).unwrap();
710
711        let stored = tokio::time::timeout(std::time::Duration::from_secs(1), async {
712            loop {
713                let stored = metadata.read().await.get_all_model_cards().pop().unwrap();
714                let DiscoveryInstance::Model { card_json, .. } = &stored else {
715                    unreachable!()
716                };
717                if card_json["runtime_config"]["taints"] == serde_json::json!(["new"]) {
718                    break stored;
719                }
720                tokio::task::yield_now().await;
721            }
722        })
723        .await
724        .expect("detached persistence did not commit local metadata");
725        let DiscoveryInstance::Model { card_json, .. } = stored else {
726            unreachable!()
727        };
728        assert_eq!(
729            card_json["runtime_config"]["taints"],
730            serde_json::json!(["new"])
731        );
732        let remote = remote.read().await.get_all_model_cards().pop().unwrap();
733        let DiscoveryInstance::Model { card_json, .. } = remote else {
734            unreachable!()
735        };
736        assert_eq!(
737            card_json["runtime_config"]["taints"],
738            serde_json::json!(["new"])
739        );
740    }
741
742    #[tokio::test]
743    async fn local_noop_reapplies_authoritative_model_taints() {
744        let local_model = model_with_taint("old");
745        let DiscoveryInstanceId::Model(id) = local_model.id() else {
746            unreachable!()
747        };
748        let mut initial = DiscoveryMetadata::new();
749        initial.register_model_card(local_model).unwrap();
750        let metadata = Arc::new(RwLock::new(initial));
751        let persisted = Arc::new(RwLock::new(None));
752        let task_persisted = persisted.clone();
753
754        let changed = update_model_taints_and_persist(
755            &metadata,
756            id,
757            HashSet::from(["old".to_string()]),
758            move |candidate| async move {
759                *task_persisted.write().await = Some(candidate.clone());
760                Ok(candidate)
761            },
762        )
763        .await
764        .unwrap();
765
766        assert!(!changed);
767        let reapplied = persisted
768            .read()
769            .await
770            .clone()
771            .expect("no-op was not persisted");
772        let DiscoveryInstance::Model { card_json, .. } =
773            reapplied.get_all_model_cards().pop().unwrap()
774        else {
775            unreachable!()
776        };
777        assert_eq!(
778            card_json["runtime_config"]["taints"],
779            serde_json::json!(["old"])
780        );
781    }
782}