1use anyhow::Result;
5use serde::Deserialize as _;
6use std::collections::{HashMap, HashSet};
7use std::sync::Arc;
8
9use super::{
10 DiscoveryInstance, DiscoveryInstanceId, DiscoveryQuery, validate_event_source_reregistration,
11};
12
13fn deserialize_null_default<'de, D, T>(deserializer: D) -> Result<T, D::Error>
27where
28 D: serde::Deserializer<'de>,
29 T: Default + serde::Deserialize<'de>,
30{
31 Ok(Option::<T>::deserialize(deserializer)?.unwrap_or_default())
32}
33
34#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
37pub struct DiscoveryMetadata {
38 #[serde(default, deserialize_with = "deserialize_null_default")]
40 endpoints: HashMap<String, DiscoveryInstance>,
41 #[serde(default, deserialize_with = "deserialize_null_default")]
43 model_cards: HashMap<String, DiscoveryInstance>,
44 #[serde(default, deserialize_with = "deserialize_null_default")]
46 event_channels: HashMap<String, DiscoveryInstance>,
47 #[serde(default, deserialize_with = "deserialize_null_default")]
49 event_sources: HashMap<String, DiscoveryInstance>,
50}
51
52impl DiscoveryMetadata {
53 pub fn new() -> Self {
55 Self {
56 endpoints: HashMap::new(),
57 model_cards: HashMap::new(),
58 event_channels: HashMap::new(),
59 event_sources: HashMap::new(),
60 }
61 }
62
63 pub fn register_endpoint(&mut self, instance: DiscoveryInstance) -> Result<()> {
65 match instance.id() {
66 DiscoveryInstanceId::Endpoint(key) => {
67 self.endpoints.insert(key.to_path(), instance);
68 Ok(())
69 }
70 DiscoveryInstanceId::Model(_) => {
71 anyhow::bail!("Cannot register non-endpoint instance as endpoint")
72 }
73 DiscoveryInstanceId::EventChannel(_) => {
74 anyhow::bail!("Cannot register EventChannel instance as endpoint")
75 }
76 DiscoveryInstanceId::EventSource(_) => {
77 anyhow::bail!("Cannot register EventSource instance as endpoint")
78 }
79 }
80 }
81
82 pub fn register_model_card(&mut self, instance: DiscoveryInstance) -> Result<()> {
84 match instance.id() {
85 DiscoveryInstanceId::Model(key) => {
86 self.model_cards.insert(key.to_path(), instance);
87 Ok(())
88 }
89 DiscoveryInstanceId::Endpoint(_) => {
90 anyhow::bail!("Cannot register non-model-card instance as model card")
91 }
92 DiscoveryInstanceId::EventChannel(_) => {
93 anyhow::bail!("Cannot register EventChannel instance as model card")
94 }
95 DiscoveryInstanceId::EventSource(_) => {
96 anyhow::bail!("Cannot register EventSource instance as model card")
97 }
98 }
99 }
100
101 pub fn unregister_endpoint(&mut self, instance: &DiscoveryInstance) -> Result<()> {
103 match instance.id() {
104 DiscoveryInstanceId::Endpoint(key) => {
105 self.endpoints.remove(&key.to_path());
106 Ok(())
107 }
108 DiscoveryInstanceId::Model(_) => {
109 anyhow::bail!("Cannot unregister non-endpoint instance as endpoint")
110 }
111 DiscoveryInstanceId::EventChannel(_) => {
112 anyhow::bail!("Cannot unregister EventChannel instance as endpoint")
113 }
114 DiscoveryInstanceId::EventSource(_) => {
115 anyhow::bail!("Cannot unregister EventSource instance as endpoint")
116 }
117 }
118 }
119
120 pub fn unregister_model_card(&mut self, instance: &DiscoveryInstance) -> Result<()> {
122 match instance.id() {
123 DiscoveryInstanceId::Model(key) => {
124 self.model_cards.remove(&key.to_path());
125 Ok(())
126 }
127 DiscoveryInstanceId::Endpoint(_) => {
128 anyhow::bail!("Cannot unregister non-model-card instance as model card")
129 }
130 DiscoveryInstanceId::EventChannel(_) => {
131 anyhow::bail!("Cannot unregister EventChannel instance as model card")
132 }
133 DiscoveryInstanceId::EventSource(_) => {
134 anyhow::bail!("Cannot unregister EventSource instance as model card")
135 }
136 }
137 }
138
139 pub fn register_event_channel(&mut self, instance: DiscoveryInstance) -> Result<()> {
141 match instance.id() {
142 DiscoveryInstanceId::EventChannel(key) => {
143 self.event_channels.insert(key.to_path(), instance);
144 Ok(())
145 }
146 DiscoveryInstanceId::Endpoint(_) => {
147 anyhow::bail!("Cannot register Endpoint instance as event channel")
148 }
149 DiscoveryInstanceId::Model(_) => {
150 anyhow::bail!("Cannot register Model instance as event channel")
151 }
152 DiscoveryInstanceId::EventSource(_) => {
153 anyhow::bail!("Cannot register EventSource instance as event channel")
154 }
155 }
156 }
157
158 pub fn unregister_event_channel(&mut self, instance: &DiscoveryInstance) -> Result<()> {
160 match instance.id() {
161 DiscoveryInstanceId::EventChannel(key) => {
162 self.event_channels.remove(&key.to_path());
163 Ok(())
164 }
165 DiscoveryInstanceId::Endpoint(_) => {
166 anyhow::bail!("Cannot unregister Endpoint instance as event channel")
167 }
168 DiscoveryInstanceId::Model(_) => {
169 anyhow::bail!("Cannot unregister Model instance as event channel")
170 }
171 DiscoveryInstanceId::EventSource(_) => {
172 anyhow::bail!("Cannot unregister EventSource instance as event channel")
173 }
174 }
175 }
176
177 pub fn register_event_source(&mut self, instance: DiscoveryInstance) -> Result<()> {
179 match instance.id() {
180 DiscoveryInstanceId::EventSource(key) => {
181 let path = key.to_path();
182 match self.event_sources.entry(path) {
183 std::collections::hash_map::Entry::Vacant(entry) => {
184 entry.insert(instance);
185 Ok(())
186 }
187 std::collections::hash_map::Entry::Occupied(entry) => {
188 validate_event_source_reregistration(entry.get(), &instance)
189 }
190 }
191 }
192 DiscoveryInstanceId::Endpoint(_) => {
193 anyhow::bail!("Cannot register Endpoint instance as event source")
194 }
195 DiscoveryInstanceId::Model(_) => {
196 anyhow::bail!("Cannot register Model instance as event source")
197 }
198 DiscoveryInstanceId::EventChannel(_) => {
199 anyhow::bail!("Cannot register EventChannel instance as event source")
200 }
201 }
202 }
203
204 pub fn unregister_event_source(&mut self, instance: &DiscoveryInstance) -> Result<()> {
206 match instance.id() {
207 DiscoveryInstanceId::EventSource(key) => {
208 self.event_sources.remove(&key.to_path());
209 Ok(())
210 }
211 DiscoveryInstanceId::Endpoint(_) => {
212 anyhow::bail!("Cannot unregister Endpoint instance as event source")
213 }
214 DiscoveryInstanceId::Model(_) => {
215 anyhow::bail!("Cannot unregister Model instance as event source")
216 }
217 DiscoveryInstanceId::EventChannel(_) => {
218 anyhow::bail!("Cannot unregister EventChannel instance as event source")
219 }
220 }
221 }
222
223 pub fn get_all_endpoints(&self) -> Vec<DiscoveryInstance> {
225 self.endpoints.values().cloned().collect()
226 }
227
228 pub fn get_all_model_cards(&self) -> Vec<DiscoveryInstance> {
230 self.model_cards.values().cloned().collect()
231 }
232
233 pub fn get_all_event_channels(&self) -> Vec<DiscoveryInstance> {
235 self.event_channels.values().cloned().collect()
236 }
237
238 pub fn get_all_event_sources(&self) -> Vec<DiscoveryInstance> {
240 self.event_sources.values().cloned().collect()
241 }
242
243 pub fn get_all(&self) -> Vec<DiscoveryInstance> {
245 self.endpoints
246 .values()
247 .chain(self.model_cards.values())
248 .chain(self.event_channels.values())
249 .chain(self.event_sources.values())
250 .cloned()
251 .collect()
252 }
253
254 pub fn filter(&self, query: &DiscoveryQuery) -> Vec<DiscoveryInstance> {
256 let all_instances = match query {
257 DiscoveryQuery::AllEndpoints
258 | DiscoveryQuery::NamespacedEndpoints { .. }
259 | DiscoveryQuery::ComponentEndpoints { .. }
260 | DiscoveryQuery::Endpoint { .. } => self.get_all_endpoints(),
261
262 DiscoveryQuery::AllModels
263 | DiscoveryQuery::NamespacedModels { .. }
264 | DiscoveryQuery::ComponentModels { .. }
265 | DiscoveryQuery::EndpointModels { .. } => self.get_all_model_cards(),
266
267 DiscoveryQuery::EventChannels(_) => self.get_all_event_channels(),
269 DiscoveryQuery::EventSources(_) => self.get_all_event_sources(),
270 };
271
272 filter_instances(all_instances, query)
273 }
274}
275
276impl Default for DiscoveryMetadata {
277 fn default() -> Self {
278 Self::new()
279 }
280}
281
282fn filter_instances(
284 instances: Vec<DiscoveryInstance>,
285 query: &DiscoveryQuery,
286) -> Vec<DiscoveryInstance> {
287 match query {
288 DiscoveryQuery::AllEndpoints | DiscoveryQuery::AllModels => instances,
289
290 DiscoveryQuery::NamespacedEndpoints { namespace } => instances
291 .into_iter()
292 .filter(|inst| match inst {
293 DiscoveryInstance::Endpoint(i) => &i.namespace == namespace,
294 _ => false,
295 })
296 .collect(),
297
298 DiscoveryQuery::ComponentEndpoints {
299 namespace,
300 component,
301 } => instances
302 .into_iter()
303 .filter(|inst| match inst {
304 DiscoveryInstance::Endpoint(i) => {
305 &i.namespace == namespace && &i.component == component
306 }
307 _ => false,
308 })
309 .collect(),
310
311 DiscoveryQuery::Endpoint {
312 namespace,
313 component,
314 endpoint,
315 } => instances
316 .into_iter()
317 .filter(|inst| match inst {
318 DiscoveryInstance::Endpoint(i) => {
319 &i.namespace == namespace
320 && &i.component == component
321 && &i.endpoint == endpoint
322 }
323 _ => false,
324 })
325 .collect(),
326
327 DiscoveryQuery::NamespacedModels { namespace } => instances
328 .into_iter()
329 .filter(|inst| match inst {
330 DiscoveryInstance::Model { namespace: ns, .. } => ns == namespace,
331 _ => false,
332 })
333 .collect(),
334
335 DiscoveryQuery::ComponentModels {
336 namespace,
337 component,
338 } => instances
339 .into_iter()
340 .filter(|inst| match inst {
341 DiscoveryInstance::Model {
342 namespace: ns,
343 component: comp,
344 ..
345 } => ns == namespace && comp == component,
346 _ => false,
347 })
348 .collect(),
349
350 DiscoveryQuery::EndpointModels {
351 namespace,
352 component,
353 endpoint,
354 } => instances
355 .into_iter()
356 .filter(|inst| match inst {
357 DiscoveryInstance::Model {
358 namespace: ns,
359 component: comp,
360 endpoint: ep,
361 ..
362 } => ns == namespace && comp == component && ep == endpoint,
363 _ => false,
364 })
365 .collect(),
366
367 DiscoveryQuery::EventChannels(query) => instances
369 .into_iter()
370 .filter(|inst| match inst {
371 DiscoveryInstance::EventChannel {
372 scope, topic: t, ..
373 } => {
374 query
375 .scope
376 .as_ref()
377 .is_none_or(|expected| expected == scope)
378 && query.topic.as_ref().is_none_or(|qt| qt == t)
379 }
380 _ => false,
381 })
382 .collect(),
383
384 DiscoveryQuery::EventSources(query) => instances
385 .into_iter()
386 .filter(|inst| match inst {
387 DiscoveryInstance::EventSource {
388 scope, topic: t, ..
389 } => {
390 query
391 .scope
392 .as_ref()
393 .is_none_or(|expected| expected == scope)
394 && query.topic.as_ref().is_none_or(|qt| qt == t)
395 }
396 _ => false,
397 })
398 .collect(),
399 }
400}
401
402#[derive(Clone, Debug)]
404pub struct MetadataSnapshot {
405 pub instances: HashMap<u64, Arc<DiscoveryMetadata>>,
407 pub generations: HashMap<u64, i64>,
410 pub sequence: u64,
412 pub timestamp: std::time::Instant,
414}
415
416impl MetadataSnapshot {
417 pub fn empty() -> Self {
418 Self {
419 instances: HashMap::new(),
420 generations: HashMap::new(),
421 sequence: 0,
422 timestamp: std::time::Instant::now(),
423 }
424 }
425
426 pub fn has_changes_from(&self, prev: &MetadataSnapshot) -> bool {
430 if self.generations == prev.generations {
431 tracing::trace!(
432 "Snapshot (seq={}): no changes, {} instances",
433 self.sequence,
434 self.instances.len()
435 );
436 return false;
437 }
438
439 let curr_ids: HashSet<u64> = self.generations.keys().copied().collect();
441 let prev_ids: HashSet<u64> = prev.generations.keys().copied().collect();
442
443 let added: Vec<_> = curr_ids
444 .difference(&prev_ids)
445 .map(|id| format!("{:x}", id))
446 .collect();
447 let removed: Vec<_> = prev_ids
448 .difference(&curr_ids)
449 .map(|id| format!("{:x}", id))
450 .collect();
451 let updated: Vec<_> = self
452 .generations
453 .iter()
454 .filter(|(k, v)| prev.generations.get(*k).is_some_and(|pv| pv != *v))
455 .map(|(k, _)| format!("{:x}", k))
456 .collect();
457
458 tracing::info!(
459 "Snapshot (seq={}): {} instances, added={:?}, removed={:?}, updated={:?}",
460 self.sequence,
461 self.instances.len(),
462 added,
463 removed,
464 updated
465 );
466
467 true
468 }
469
470 pub fn filter(&self, query: &DiscoveryQuery) -> Vec<DiscoveryInstance> {
472 self.instances
473 .values()
474 .flat_map(|metadata| metadata.filter(query))
475 .collect()
476 }
477}
478
479#[cfg(test)]
480mod tests {
481 use super::*;
482 use crate::component::{Instance, TransportType};
483 use crate::discovery::{EventChannelQuery, EventSourceQuery};
484
485 #[test]
486 fn test_metadata_serde() {
487 let mut metadata = DiscoveryMetadata::new();
488
489 let instance = DiscoveryInstance::Endpoint(Instance {
491 namespace: "test".to_string(),
492 component: "comp1".to_string(),
493 endpoint: "ep1".to_string(),
494 instance_id: 123,
495 transport: TransportType::Nats("nats://localhost:4222".to_string()),
496 device_type: None,
497 request_plane_codec: None,
498 });
499
500 metadata.register_endpoint(instance).unwrap();
501
502 let json = serde_json::to_string(&metadata).unwrap();
504
505 let deserialized: DiscoveryMetadata = serde_json::from_str(&json).unwrap();
507
508 assert_eq!(deserialized.endpoints.len(), 1);
509 assert_eq!(deserialized.model_cards.len(), 0);
510 }
511
512 #[tokio::test]
513 async fn test_metadata_accessors() {
514 let mut metadata = DiscoveryMetadata::new();
515
516 for i in 0..3 {
518 let instance = DiscoveryInstance::Endpoint(Instance {
519 namespace: "test".to_string(),
520 component: "comp1".to_string(),
521 endpoint: format!("ep{}", i),
522 instance_id: i,
523 transport: TransportType::Nats("nats://localhost:4222".to_string()),
524 device_type: None,
525 request_plane_codec: None,
526 });
527 metadata.register_endpoint(instance).unwrap();
528 }
529
530 for i in 0..2 {
532 let instance = DiscoveryInstance::Model {
533 namespace: "test".to_string(),
534 component: "comp1".to_string(),
535 endpoint: format!("ep{}", i),
536 instance_id: i,
537 card_json: serde_json::json!({"model": "test"}),
538 model_suffix: None,
539 };
540 metadata.register_model_card(instance).unwrap();
541 }
542
543 assert_eq!(metadata.get_all_endpoints().len(), 3);
544 assert_eq!(metadata.get_all_model_cards().len(), 2);
545 assert_eq!(metadata.get_all().len(), 5);
546 }
547
548 #[test]
549 fn event_source_registration_filters_and_removes_exact_incarnation() {
550 use crate::discovery::EventScope;
551 use crate::protocols::EndpointId;
552
553 let mut metadata = DiscoveryMetadata::new();
554 let endpoint = EndpointId {
555 namespace: "test".to_string(),
556 component: "worker".to_string(),
557 name: "decode".to_string(),
558 };
559 let source = |publisher_id| DiscoveryInstance::EventSource {
560 scope: EventScope::Endpoint {
561 endpoint: endpoint.clone(),
562 },
563 topic: "kv-events".to_string(),
564 publisher_id,
565 metadata: serde_json::json!({"worker_id": 7, "dp_rank": 0}),
566 };
567
568 let old = source(100);
569 let current = source(205);
570 metadata.register_event_source(old.clone()).unwrap();
571 metadata.register_event_source(old.clone()).unwrap();
572 let mut conflicting = old.clone();
573 let DiscoveryInstance::EventSource {
574 metadata: descriptor,
575 ..
576 } = &mut conflicting
577 else {
578 unreachable!()
579 };
580 *descriptor = serde_json::json!({"worker_id": 7, "dp_rank": 0, "changed": true});
581 assert!(
582 metadata
583 .register_event_source(conflicting)
584 .unwrap_err()
585 .to_string()
586 .contains("cannot change its descriptor")
587 );
588 metadata.register_event_source(current.clone()).unwrap();
589
590 let query =
591 DiscoveryQuery::EventSources(EventSourceQuery::endpoint_topic(endpoint, "kv-events"));
592 assert_eq!(metadata.filter(&query).len(), 2);
593
594 metadata.unregister_event_source(&old).unwrap();
595 assert_eq!(metadata.filter(&query), vec![current]);
596 }
597
598 #[test]
599 fn event_channel_queries_match_exact_scope() {
600 use crate::discovery::{EventScope, EventTransport};
601 use crate::protocols::EndpointId;
602
603 let mut metadata = DiscoveryMetadata::new();
604 let component_scope = EventScope::Component {
605 namespace: "test".to_string(),
606 component: "worker".to_string(),
607 };
608 let endpoint_a = EndpointId {
609 namespace: "test".to_string(),
610 component: "worker".to_string(),
611 name: "a".to_string(),
612 };
613 let endpoint_b = EndpointId {
614 name: "b".to_string(),
615 ..endpoint_a.clone()
616 };
617
618 for (instance_id, scope) in [
619 (1, component_scope),
620 (
621 2,
622 EventScope::Endpoint {
623 endpoint: endpoint_a.clone(),
624 },
625 ),
626 (
627 3,
628 EventScope::Endpoint {
629 endpoint: endpoint_b.clone(),
630 },
631 ),
632 ] {
633 metadata
634 .register_event_channel(DiscoveryInstance::EventChannel {
635 scope,
636 topic: "kv-events".to_string(),
637 instance_id,
638 transport: EventTransport::zmq(format!("tcp://localhost:{instance_id}")),
639 })
640 .unwrap();
641 }
642
643 assert_eq!(metadata.get_all_event_channels().len(), 3);
644 assert_eq!(metadata.get_all().len(), 3);
645 assert_eq!(
646 metadata
647 .filter(&DiscoveryQuery::EventChannels(EventChannelQuery::all()))
648 .len(),
649 3
650 );
651
652 let component = metadata.filter(&DiscoveryQuery::EventChannels(EventChannelQuery::topic(
653 "test",
654 "worker",
655 "kv-events",
656 )));
657 assert_eq!(component.len(), 1);
658
659 let endpoint_a_instances = metadata.filter(&DiscoveryQuery::EventChannels(
660 EventChannelQuery::endpoint_topic(endpoint_a, "kv-events"),
661 ));
662 assert_eq!(endpoint_a_instances.len(), 1);
663 assert_eq!(endpoint_a_instances[0].instance_id(), 2);
664
665 let endpoint_b_instances = metadata.filter(&DiscoveryQuery::EventChannels(
666 EventChannelQuery::endpoint_topic(endpoint_b, "kv-events"),
667 ));
668 assert_eq!(endpoint_b_instances.len(), 1);
669 assert_eq!(endpoint_b_instances[0].instance_id(), 3);
670
671 metadata
672 .unregister_event_channel(&endpoint_a_instances[0])
673 .unwrap();
674 assert_eq!(metadata.get_all_event_channels().len(), 2);
675 assert!(
676 metadata
677 .filter(&DiscoveryQuery::EventChannels(
678 EventChannelQuery::component("other", "worker")
679 ))
680 .is_empty()
681 );
682 }
683
684 #[tokio::test]
685 async fn test_mixed_instances() {
686 use crate::discovery::{EventScope, EventTransport};
687
688 let mut metadata = DiscoveryMetadata::new();
689
690 let endpoint = DiscoveryInstance::Endpoint(Instance {
692 namespace: "test".to_string(),
693 component: "comp1".to_string(),
694 endpoint: "ep1".to_string(),
695 instance_id: 1,
696 transport: TransportType::Nats("nats://localhost:4222".to_string()),
697 device_type: None,
698 request_plane_codec: None,
699 });
700 metadata.register_endpoint(endpoint).unwrap();
701
702 let model = DiscoveryInstance::Model {
703 namespace: "test".to_string(),
704 component: "comp1".to_string(),
705 endpoint: "ep1".to_string(),
706 instance_id: 2,
707 card_json: serde_json::json!({"model": "test"}),
708 model_suffix: None,
709 };
710 metadata.register_model_card(model).unwrap();
711
712 let event_channel = DiscoveryInstance::EventChannel {
713 scope: EventScope::Component {
714 namespace: "test".to_string(),
715 component: "comp1".to_string(),
716 },
717 topic: "test-topic".to_string(),
718 instance_id: 3,
719 transport: EventTransport::zmq("tcp://localhost:5000"),
720 };
721 metadata.register_event_channel(event_channel).unwrap();
722
723 assert_eq!(metadata.get_all().len(), 3);
725 assert_eq!(metadata.get_all_endpoints().len(), 1);
726 assert_eq!(metadata.get_all_model_cards().len(), 1);
727 assert_eq!(metadata.get_all_event_channels().len(), 1);
728 }
729}