Skip to main content

dynamo_runtime/discovery/
metadata.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use anyhow::Result;
5use serde::Deserialize as _;
6use std::collections::{HashMap, HashSet};
7use std::sync::Arc;
8
9use super::{
10    DiscoveryInstance, DiscoveryInstanceId, DiscoveryQuery, ModelCardInstanceId,
11    model_with_updated_taints, validate_event_source_reregistration, validate_model_reregistration,
12};
13
14/// Deserializes a JSON `null` or missing field as `T::default()`.
15///
16/// Kubernetes Server-Side Apply with `schema = "disabled"` can write an empty
17/// object `{}` as `null` for nested free-form fields. Without this helper, the
18/// daemon fails to deserialize the `DynamoWorkerMetadata` CR, and the worker is
19/// excluded from the `MetadataSnapshot` (i.e. invisible to service discovery),
20/// causing `KubeDiscoveryClient::list` to return 0 instances and all inference
21/// requests to 404. One concrete example is vLLM elastic EP scaling:
22/// `scale_elastic_ep` reinitializes event plane sockets, which triggers
23/// `unregister_event_channel()`, leaving `event_channels` as an empty map `{}`.
24/// SSA then writes it back as `null`, breaking deserialization until this helper
25/// treats `null` as an empty map. The issue applies to any event plane
26/// implementation, not only a specific transport.
27fn deserialize_null_default<'de, D, T>(deserializer: D) -> Result<T, D::Error>
28where
29    D: serde::Deserializer<'de>,
30    T: Default + serde::Deserialize<'de>,
31{
32    Ok(Option::<T>::deserialize(deserializer)?.unwrap_or_default())
33}
34
35/// Metadata stored on each pod and exposed via HTTP endpoint
36/// This struct holds all discovery registrations for this pod instance
37#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
38pub struct DiscoveryMetadata {
39    /// Registered endpoint instances (key: path string from EndpointInstanceId::to_path())
40    #[serde(default, deserialize_with = "deserialize_null_default")]
41    endpoints: HashMap<String, DiscoveryInstance>,
42    /// Registered model card instances (key: path string from ModelCardInstanceId::to_path())
43    #[serde(default, deserialize_with = "deserialize_null_default")]
44    model_cards: HashMap<String, DiscoveryInstance>,
45    /// Registered event channel instances (key: path string from EventChannelInstanceId::to_path())
46    #[serde(default, deserialize_with = "deserialize_null_default")]
47    event_channels: HashMap<String, DiscoveryInstance>,
48    /// Registered event source instances (key: path string from EventSourceInstanceId::to_path())
49    #[serde(default, deserialize_with = "deserialize_null_default")]
50    event_sources: HashMap<String, DiscoveryInstance>,
51}
52
53impl DiscoveryMetadata {
54    /// Create a new empty metadata store
55    pub fn new() -> Self {
56        Self {
57            endpoints: HashMap::new(),
58            model_cards: HashMap::new(),
59            event_channels: HashMap::new(),
60            event_sources: HashMap::new(),
61        }
62    }
63
64    /// Register an endpoint instance
65    pub fn register_endpoint(&mut self, instance: DiscoveryInstance) -> Result<()> {
66        match instance.id() {
67            DiscoveryInstanceId::Endpoint(key) => {
68                self.endpoints.insert(key.to_path(), instance);
69                Ok(())
70            }
71            DiscoveryInstanceId::Model(_) => {
72                anyhow::bail!("Cannot register non-endpoint instance as endpoint")
73            }
74            DiscoveryInstanceId::EventChannel(_) => {
75                anyhow::bail!("Cannot register EventChannel instance as endpoint")
76            }
77            DiscoveryInstanceId::EventSource(_) => {
78                anyhow::bail!("Cannot register EventSource instance as endpoint")
79            }
80        }
81    }
82
83    /// Register a model card instance
84    pub fn register_model_card(
85        &mut self,
86        instance: DiscoveryInstance,
87    ) -> Result<DiscoveryInstance> {
88        match instance.id() {
89            DiscoveryInstanceId::Model(key) => match self.model_cards.entry(key.to_path()) {
90                std::collections::hash_map::Entry::Vacant(entry) => {
91                    entry.insert(instance.clone());
92                    Ok(instance)
93                }
94                std::collections::hash_map::Entry::Occupied(entry) => {
95                    validate_model_reregistration(entry.get(), &instance)?;
96                    Ok(entry.get().clone())
97                }
98            },
99            DiscoveryInstanceId::Endpoint(_) => {
100                anyhow::bail!("Cannot register non-model-card instance as model card")
101            }
102            DiscoveryInstanceId::EventChannel(_) => {
103                anyhow::bail!("Cannot register EventChannel instance as model card")
104            }
105            DiscoveryInstanceId::EventSource(_) => {
106                anyhow::bail!("Cannot register EventSource instance as model card")
107            }
108        }
109    }
110
111    /// Update one authoritative model card while the caller holds the metadata lock.
112    pub fn update_model_taints(
113        &mut self,
114        id: &ModelCardInstanceId,
115        taints: HashSet<String>,
116    ) -> Result<bool> {
117        let path = id.to_path();
118        let existing = self
119            .model_cards
120            .get_mut(&path)
121            .ok_or_else(|| anyhow::anyhow!("model discovery record {path} is not registered"))?;
122        let candidate = model_with_updated_taints(existing, taints)?;
123        if candidate == *existing {
124            return Ok(false);
125        }
126        *existing = candidate;
127        Ok(true)
128    }
129
130    /// Unregister an endpoint instance
131    pub fn unregister_endpoint(&mut self, instance: &DiscoveryInstance) -> Result<()> {
132        match instance.id() {
133            DiscoveryInstanceId::Endpoint(key) => {
134                self.endpoints.remove(&key.to_path());
135                Ok(())
136            }
137            DiscoveryInstanceId::Model(_) => {
138                anyhow::bail!("Cannot unregister non-endpoint instance as endpoint")
139            }
140            DiscoveryInstanceId::EventChannel(_) => {
141                anyhow::bail!("Cannot unregister EventChannel instance as endpoint")
142            }
143            DiscoveryInstanceId::EventSource(_) => {
144                anyhow::bail!("Cannot unregister EventSource instance as endpoint")
145            }
146        }
147    }
148
149    /// Unregister a model card instance
150    pub fn unregister_model_card(&mut self, instance: &DiscoveryInstance) -> Result<()> {
151        match instance.id() {
152            DiscoveryInstanceId::Model(key) => {
153                self.model_cards.remove(&key.to_path());
154                Ok(())
155            }
156            DiscoveryInstanceId::Endpoint(_) => {
157                anyhow::bail!("Cannot unregister non-model-card instance as model card")
158            }
159            DiscoveryInstanceId::EventChannel(_) => {
160                anyhow::bail!("Cannot unregister EventChannel instance as model card")
161            }
162            DiscoveryInstanceId::EventSource(_) => {
163                anyhow::bail!("Cannot unregister EventSource instance as model card")
164            }
165        }
166    }
167
168    /// Register an event channel instance
169    pub fn register_event_channel(&mut self, instance: DiscoveryInstance) -> Result<()> {
170        match instance.id() {
171            DiscoveryInstanceId::EventChannel(key) => {
172                self.event_channels.insert(key.to_path(), instance);
173                Ok(())
174            }
175            DiscoveryInstanceId::Endpoint(_) => {
176                anyhow::bail!("Cannot register Endpoint instance as event channel")
177            }
178            DiscoveryInstanceId::Model(_) => {
179                anyhow::bail!("Cannot register Model instance as event channel")
180            }
181            DiscoveryInstanceId::EventSource(_) => {
182                anyhow::bail!("Cannot register EventSource instance as event channel")
183            }
184        }
185    }
186
187    /// Unregister an event channel instance
188    pub fn unregister_event_channel(&mut self, instance: &DiscoveryInstance) -> Result<()> {
189        match instance.id() {
190            DiscoveryInstanceId::EventChannel(key) => {
191                self.event_channels.remove(&key.to_path());
192                Ok(())
193            }
194            DiscoveryInstanceId::Endpoint(_) => {
195                anyhow::bail!("Cannot unregister Endpoint instance as event channel")
196            }
197            DiscoveryInstanceId::Model(_) => {
198                anyhow::bail!("Cannot unregister Model instance as event channel")
199            }
200            DiscoveryInstanceId::EventSource(_) => {
201                anyhow::bail!("Cannot unregister EventSource instance as event channel")
202            }
203        }
204    }
205
206    /// Register a semantic event source instance.
207    pub fn register_event_source(&mut self, instance: DiscoveryInstance) -> Result<()> {
208        match instance.id() {
209            DiscoveryInstanceId::EventSource(key) => {
210                let path = key.to_path();
211                match self.event_sources.entry(path) {
212                    std::collections::hash_map::Entry::Vacant(entry) => {
213                        entry.insert(instance);
214                        Ok(())
215                    }
216                    std::collections::hash_map::Entry::Occupied(entry) => {
217                        validate_event_source_reregistration(entry.get(), &instance)
218                    }
219                }
220            }
221            DiscoveryInstanceId::Endpoint(_) => {
222                anyhow::bail!("Cannot register Endpoint instance as event source")
223            }
224            DiscoveryInstanceId::Model(_) => {
225                anyhow::bail!("Cannot register Model instance as event source")
226            }
227            DiscoveryInstanceId::EventChannel(_) => {
228                anyhow::bail!("Cannot register EventChannel instance as event source")
229            }
230        }
231    }
232
233    /// Unregister one exact semantic event source incarnation.
234    pub fn unregister_event_source(&mut self, instance: &DiscoveryInstance) -> Result<()> {
235        match instance.id() {
236            DiscoveryInstanceId::EventSource(key) => {
237                self.event_sources.remove(&key.to_path());
238                Ok(())
239            }
240            DiscoveryInstanceId::Endpoint(_) => {
241                anyhow::bail!("Cannot unregister Endpoint instance as event source")
242            }
243            DiscoveryInstanceId::Model(_) => {
244                anyhow::bail!("Cannot unregister Model instance as event source")
245            }
246            DiscoveryInstanceId::EventChannel(_) => {
247                anyhow::bail!("Cannot unregister EventChannel instance as event source")
248            }
249        }
250    }
251
252    /// Get all registered endpoints
253    pub fn get_all_endpoints(&self) -> Vec<DiscoveryInstance> {
254        self.endpoints.values().cloned().collect()
255    }
256
257    /// Get all registered model cards
258    pub fn get_all_model_cards(&self) -> Vec<DiscoveryInstance> {
259        self.model_cards.values().cloned().collect()
260    }
261
262    /// Get all registered event channels
263    pub fn get_all_event_channels(&self) -> Vec<DiscoveryInstance> {
264        self.event_channels.values().cloned().collect()
265    }
266
267    /// Get all registered semantic event sources.
268    pub fn get_all_event_sources(&self) -> Vec<DiscoveryInstance> {
269        self.event_sources.values().cloned().collect()
270    }
271
272    /// Get all registered instances (endpoints, model cards, and event channels)
273    pub fn get_all(&self) -> Vec<DiscoveryInstance> {
274        self.endpoints
275            .values()
276            .chain(self.model_cards.values())
277            .chain(self.event_channels.values())
278            .chain(self.event_sources.values())
279            .cloned()
280            .collect()
281    }
282
283    /// Filter this metadata by query
284    pub fn filter(&self, query: &DiscoveryQuery) -> Vec<DiscoveryInstance> {
285        let all_instances = match query {
286            DiscoveryQuery::AllEndpoints
287            | DiscoveryQuery::NamespacedEndpoints { .. }
288            | DiscoveryQuery::ComponentEndpoints { .. }
289            | DiscoveryQuery::Endpoint { .. } => self.get_all_endpoints(),
290
291            DiscoveryQuery::AllModels
292            | DiscoveryQuery::NamespacedModels { .. }
293            | DiscoveryQuery::ComponentModels { .. }
294            | DiscoveryQuery::EndpointModels { .. } => self.get_all_model_cards(),
295
296            // EventChannel queries now return actual event channels
297            DiscoveryQuery::EventChannels(_) => self.get_all_event_channels(),
298            DiscoveryQuery::EventSources(_) => self.get_all_event_sources(),
299        };
300
301        filter_instances(all_instances, query)
302    }
303}
304
305impl Default for DiscoveryMetadata {
306    fn default() -> Self {
307        Self::new()
308    }
309}
310
311/// Filter instances by query predicate
312fn filter_instances(
313    instances: Vec<DiscoveryInstance>,
314    query: &DiscoveryQuery,
315) -> Vec<DiscoveryInstance> {
316    match query {
317        DiscoveryQuery::AllEndpoints | DiscoveryQuery::AllModels => instances,
318
319        DiscoveryQuery::NamespacedEndpoints { namespace } => instances
320            .into_iter()
321            .filter(|inst| match inst {
322                DiscoveryInstance::Endpoint(i) => &i.namespace == namespace,
323                _ => false,
324            })
325            .collect(),
326
327        DiscoveryQuery::ComponentEndpoints {
328            namespace,
329            component,
330        } => instances
331            .into_iter()
332            .filter(|inst| match inst {
333                DiscoveryInstance::Endpoint(i) => {
334                    &i.namespace == namespace && &i.component == component
335                }
336                _ => false,
337            })
338            .collect(),
339
340        DiscoveryQuery::Endpoint {
341            namespace,
342            component,
343            endpoint,
344        } => instances
345            .into_iter()
346            .filter(|inst| match inst {
347                DiscoveryInstance::Endpoint(i) => {
348                    &i.namespace == namespace
349                        && &i.component == component
350                        && &i.endpoint == endpoint
351                }
352                _ => false,
353            })
354            .collect(),
355
356        DiscoveryQuery::NamespacedModels { namespace } => instances
357            .into_iter()
358            .filter(|inst| match inst {
359                DiscoveryInstance::Model { namespace: ns, .. } => ns == namespace,
360                _ => false,
361            })
362            .collect(),
363
364        DiscoveryQuery::ComponentModels {
365            namespace,
366            component,
367        } => instances
368            .into_iter()
369            .filter(|inst| match inst {
370                DiscoveryInstance::Model {
371                    namespace: ns,
372                    component: comp,
373                    ..
374                } => ns == namespace && comp == component,
375                _ => false,
376            })
377            .collect(),
378
379        DiscoveryQuery::EndpointModels {
380            namespace,
381            component,
382            endpoint,
383        } => instances
384            .into_iter()
385            .filter(|inst| match inst {
386                DiscoveryInstance::Model {
387                    namespace: ns,
388                    component: comp,
389                    endpoint: ep,
390                    ..
391                } => ns == namespace && comp == component && ep == endpoint,
392                _ => false,
393            })
394            .collect(),
395
396        // EventChannel queries - unified filtering with optional scope filters
397        DiscoveryQuery::EventChannels(query) => instances
398            .into_iter()
399            .filter(|inst| match inst {
400                DiscoveryInstance::EventChannel {
401                    scope, topic: t, ..
402                } => {
403                    query
404                        .scope
405                        .as_ref()
406                        .is_none_or(|expected| expected == scope)
407                        && query.topic.as_ref().is_none_or(|qt| qt == t)
408                }
409                _ => false,
410            })
411            .collect(),
412
413        DiscoveryQuery::EventSources(query) => instances
414            .into_iter()
415            .filter(|inst| match inst {
416                DiscoveryInstance::EventSource {
417                    scope, topic: t, ..
418                } => {
419                    query
420                        .scope
421                        .as_ref()
422                        .is_none_or(|expected| expected == scope)
423                        && query.topic.as_ref().is_none_or(|qt| qt == t)
424                }
425                _ => false,
426            })
427            .collect(),
428    }
429}
430
431/// Snapshot of all discovered instances and their metadata
432#[derive(Clone, Debug)]
433pub struct MetadataSnapshot {
434    /// Map of instance_id -> metadata
435    pub instances: HashMap<u64, Arc<DiscoveryMetadata>>,
436    /// Map of instance_id -> CR generation for change detection
437    /// Keys match `instances` keys exactly - only ready pods with CRs are included
438    pub generations: HashMap<u64, i64>,
439    /// Sequence number for debugging
440    pub sequence: u64,
441    /// Timestamp for observability
442    pub timestamp: std::time::Instant,
443}
444
445impl MetadataSnapshot {
446    pub fn empty() -> Self {
447        Self {
448            instances: HashMap::new(),
449            generations: HashMap::new(),
450            sequence: 0,
451            timestamp: std::time::Instant::now(),
452        }
453    }
454
455    /// Compare with previous snapshot and return true if changed.
456    /// Logs diagnostic info about what changed.
457    /// This is done on the basis of the generation of the DynamoWorkerMetadata CRs that are owned by ready workers
458    pub fn has_changes_from(&self, prev: &MetadataSnapshot) -> bool {
459        if self.generations == prev.generations {
460            tracing::trace!(
461                "Snapshot (seq={}): no changes, {} instances",
462                self.sequence,
463                self.instances.len()
464            );
465            return false;
466        }
467
468        // Compute diff for logging
469        let curr_ids: HashSet<u64> = self.generations.keys().copied().collect();
470        let prev_ids: HashSet<u64> = prev.generations.keys().copied().collect();
471
472        let added: Vec<_> = curr_ids
473            .difference(&prev_ids)
474            .map(|id| format!("{:x}", id))
475            .collect();
476        let removed: Vec<_> = prev_ids
477            .difference(&curr_ids)
478            .map(|id| format!("{:x}", id))
479            .collect();
480        let updated: Vec<_> = self
481            .generations
482            .iter()
483            .filter(|(k, v)| prev.generations.get(*k).is_some_and(|pv| pv != *v))
484            .map(|(k, _)| format!("{:x}", k))
485            .collect();
486
487        tracing::info!(
488            "Snapshot (seq={}): {} instances, added={:?}, removed={:?}, updated={:?}",
489            self.sequence,
490            self.instances.len(),
491            added,
492            removed,
493            updated
494        );
495
496        true
497    }
498
499    /// Filter all instances in the snapshot by query
500    pub fn filter(&self, query: &DiscoveryQuery) -> Vec<DiscoveryInstance> {
501        self.instances
502            .values()
503            .flat_map(|metadata| metadata.filter(query))
504            .collect()
505    }
506}
507
508#[cfg(test)]
509mod tests {
510    use super::*;
511    use crate::component::{Instance, TransportType};
512    use crate::discovery::{EventChannelQuery, EventSourceQuery};
513
514    #[test]
515    fn authoritative_model_taint_updates_do_not_read_stale_snapshots() {
516        let mut metadata = DiscoveryMetadata::new();
517        let model = DiscoveryInstance::Model {
518            namespace: "ns".to_string(),
519            component: "worker".to_string(),
520            endpoint: "generate".to_string(),
521            instance_id: 7,
522            card_json: serde_json::json!({
523                "runtime_config": {"taints": ["a"]}
524            }),
525            model_suffix: None,
526        };
527        let DiscoveryInstanceId::Model(id) = model.id() else {
528            unreachable!()
529        };
530        metadata.register_model_card(model.clone()).unwrap();
531
532        assert!(
533            metadata
534                .update_model_taints(&id, HashSet::from(["b".to_string()]))
535                .unwrap()
536        );
537        assert!(
538            metadata
539                .update_model_taints(&id, HashSet::from(["a".to_string()]))
540                .unwrap()
541        );
542
543        let stored = metadata.get_all_model_cards().pop().unwrap();
544        let DiscoveryInstance::Model { card_json, .. } = stored else {
545            unreachable!()
546        };
547        assert_eq!(
548            card_json["runtime_config"]["taints"],
549            serde_json::json!(["a"])
550        );
551
552        metadata.unregister_model_card(&model).unwrap();
553        let error = metadata
554            .update_model_taints(&id, HashSet::from(["b".to_string()]))
555            .unwrap_err();
556        assert!(error.to_string().contains("is not registered"));
557    }
558
559    #[test]
560    fn same_id_model_registration_preserves_authoritative_taints() {
561        let mut metadata = DiscoveryMetadata::new();
562        let model = DiscoveryInstance::Model {
563            namespace: "ns".to_string(),
564            component: "worker".to_string(),
565            endpoint: "generate".to_string(),
566            instance_id: 7,
567            card_json: serde_json::json!({
568                "display_name": "model",
569                "runtime_config": {"taints": ["initial"]}
570            }),
571            model_suffix: None,
572        };
573        let DiscoveryInstanceId::Model(id) = model.id() else {
574            unreachable!()
575        };
576        metadata.register_model_card(model.clone()).unwrap();
577        metadata
578            .update_model_taints(&id, HashSet::from(["updated".to_string()]))
579            .unwrap();
580
581        let replayed = metadata.register_model_card(model).unwrap();
582
583        let DiscoveryInstance::Model { card_json, .. } = replayed else {
584            unreachable!()
585        };
586        assert_eq!(
587            card_json["runtime_config"]["taints"],
588            serde_json::json!(["updated"])
589        );
590        assert_eq!(metadata.get_all_model_cards().len(), 1);
591    }
592
593    #[test]
594    fn test_metadata_serde() {
595        let mut metadata = DiscoveryMetadata::new();
596
597        // Add an endpoint
598        let instance = DiscoveryInstance::Endpoint(Instance {
599            namespace: "test".to_string(),
600            component: "comp1".to_string(),
601            endpoint: "ep1".to_string(),
602            instance_id: 123,
603            transport: TransportType::Nats("nats://localhost:4222".to_string()),
604            device_type: None,
605            request_plane_codec: None,
606        });
607
608        metadata.register_endpoint(instance).unwrap();
609
610        // Serialize
611        let json = serde_json::to_string(&metadata).unwrap();
612
613        // Deserialize
614        let deserialized: DiscoveryMetadata = serde_json::from_str(&json).unwrap();
615
616        assert_eq!(deserialized.endpoints.len(), 1);
617        assert_eq!(deserialized.model_cards.len(), 0);
618    }
619
620    #[tokio::test]
621    async fn test_metadata_accessors() {
622        let mut metadata = DiscoveryMetadata::new();
623
624        // Register endpoints
625        for i in 0..3 {
626            let instance = DiscoveryInstance::Endpoint(Instance {
627                namespace: "test".to_string(),
628                component: "comp1".to_string(),
629                endpoint: format!("ep{}", i),
630                instance_id: i,
631                transport: TransportType::Nats("nats://localhost:4222".to_string()),
632                device_type: None,
633                request_plane_codec: None,
634            });
635            metadata.register_endpoint(instance).unwrap();
636        }
637
638        // Register model cards
639        for i in 0..2 {
640            let instance = DiscoveryInstance::Model {
641                namespace: "test".to_string(),
642                component: "comp1".to_string(),
643                endpoint: format!("ep{}", i),
644                instance_id: i,
645                card_json: serde_json::json!({"model": "test"}),
646                model_suffix: None,
647            };
648            metadata.register_model_card(instance).unwrap();
649        }
650
651        assert_eq!(metadata.get_all_endpoints().len(), 3);
652        assert_eq!(metadata.get_all_model_cards().len(), 2);
653        assert_eq!(metadata.get_all().len(), 5);
654    }
655
656    #[test]
657    fn event_source_registration_filters_and_removes_exact_incarnation() {
658        use crate::discovery::EventScope;
659        use crate::protocols::EndpointId;
660
661        let mut metadata = DiscoveryMetadata::new();
662        let endpoint = EndpointId {
663            namespace: "test".to_string(),
664            component: "worker".to_string(),
665            name: "decode".to_string(),
666        };
667        let source = |publisher_id| DiscoveryInstance::EventSource {
668            scope: EventScope::Endpoint {
669                endpoint: endpoint.clone(),
670            },
671            topic: "kv-events".to_string(),
672            publisher_id,
673            metadata: serde_json::json!({"worker_id": 7, "dp_rank": 0}),
674        };
675
676        let old = source(100);
677        let current = source(205);
678        metadata.register_event_source(old.clone()).unwrap();
679        metadata.register_event_source(old.clone()).unwrap();
680        let mut conflicting = old.clone();
681        let DiscoveryInstance::EventSource {
682            metadata: descriptor,
683            ..
684        } = &mut conflicting
685        else {
686            unreachable!()
687        };
688        *descriptor = serde_json::json!({"worker_id": 7, "dp_rank": 0, "changed": true});
689        assert!(
690            metadata
691                .register_event_source(conflicting)
692                .unwrap_err()
693                .to_string()
694                .contains("cannot change its descriptor")
695        );
696        metadata.register_event_source(current.clone()).unwrap();
697
698        let query =
699            DiscoveryQuery::EventSources(EventSourceQuery::endpoint_topic(endpoint, "kv-events"));
700        assert_eq!(metadata.filter(&query).len(), 2);
701
702        metadata.unregister_event_source(&old).unwrap();
703        assert_eq!(metadata.filter(&query), vec![current]);
704    }
705
706    #[test]
707    fn event_channel_queries_match_exact_scope() {
708        use crate::discovery::{EventScope, EventTransport};
709        use crate::protocols::EndpointId;
710
711        let mut metadata = DiscoveryMetadata::new();
712        let component_scope = EventScope::Component {
713            namespace: "test".to_string(),
714            component: "worker".to_string(),
715        };
716        let endpoint_a = EndpointId {
717            namespace: "test".to_string(),
718            component: "worker".to_string(),
719            name: "a".to_string(),
720        };
721        let endpoint_b = EndpointId {
722            name: "b".to_string(),
723            ..endpoint_a.clone()
724        };
725
726        for (instance_id, scope) in [
727            (1, component_scope),
728            (
729                2,
730                EventScope::Endpoint {
731                    endpoint: endpoint_a.clone(),
732                },
733            ),
734            (
735                3,
736                EventScope::Endpoint {
737                    endpoint: endpoint_b.clone(),
738                },
739            ),
740        ] {
741            metadata
742                .register_event_channel(DiscoveryInstance::EventChannel {
743                    scope,
744                    topic: "kv-events".to_string(),
745                    instance_id,
746                    transport: EventTransport::zmq(format!("tcp://localhost:{instance_id}")),
747                })
748                .unwrap();
749        }
750
751        assert_eq!(metadata.get_all_event_channels().len(), 3);
752        assert_eq!(metadata.get_all().len(), 3);
753        assert_eq!(
754            metadata
755                .filter(&DiscoveryQuery::EventChannels(EventChannelQuery::all()))
756                .len(),
757            3
758        );
759
760        let component = metadata.filter(&DiscoveryQuery::EventChannels(EventChannelQuery::topic(
761            "test",
762            "worker",
763            "kv-events",
764        )));
765        assert_eq!(component.len(), 1);
766
767        let endpoint_a_instances = metadata.filter(&DiscoveryQuery::EventChannels(
768            EventChannelQuery::endpoint_topic(endpoint_a, "kv-events"),
769        ));
770        assert_eq!(endpoint_a_instances.len(), 1);
771        assert_eq!(endpoint_a_instances[0].instance_id(), 2);
772
773        let endpoint_b_instances = metadata.filter(&DiscoveryQuery::EventChannels(
774            EventChannelQuery::endpoint_topic(endpoint_b, "kv-events"),
775        ));
776        assert_eq!(endpoint_b_instances.len(), 1);
777        assert_eq!(endpoint_b_instances[0].instance_id(), 3);
778
779        metadata
780            .unregister_event_channel(&endpoint_a_instances[0])
781            .unwrap();
782        assert_eq!(metadata.get_all_event_channels().len(), 2);
783        assert!(
784            metadata
785                .filter(&DiscoveryQuery::EventChannels(
786                    EventChannelQuery::component("other", "worker")
787                ))
788                .is_empty()
789        );
790    }
791
792    #[tokio::test]
793    async fn test_mixed_instances() {
794        use crate::discovery::{EventScope, EventTransport};
795
796        let mut metadata = DiscoveryMetadata::new();
797
798        // Register one of each type
799        let endpoint = DiscoveryInstance::Endpoint(Instance {
800            namespace: "test".to_string(),
801            component: "comp1".to_string(),
802            endpoint: "ep1".to_string(),
803            instance_id: 1,
804            transport: TransportType::Nats("nats://localhost:4222".to_string()),
805            device_type: None,
806            request_plane_codec: None,
807        });
808        metadata.register_endpoint(endpoint).unwrap();
809
810        let model = DiscoveryInstance::Model {
811            namespace: "test".to_string(),
812            component: "comp1".to_string(),
813            endpoint: "ep1".to_string(),
814            instance_id: 2,
815            card_json: serde_json::json!({"model": "test"}),
816            model_suffix: None,
817        };
818        metadata.register_model_card(model).unwrap();
819
820        let event_channel = DiscoveryInstance::EventChannel {
821            scope: EventScope::Component {
822                namespace: "test".to_string(),
823                component: "comp1".to_string(),
824            },
825            topic: "test-topic".to_string(),
826            instance_id: 3,
827            transport: EventTransport::zmq("tcp://localhost:5000"),
828        };
829        metadata.register_event_channel(event_channel).unwrap();
830
831        // Verify get_all returns all three
832        assert_eq!(metadata.get_all().len(), 3);
833        assert_eq!(metadata.get_all_endpoints().len(), 1);
834        assert_eq!(metadata.get_all_model_cards().len(), 1);
835        assert_eq!(metadata.get_all_event_channels().len(), 1);
836    }
837}