Skip to main content

dynamo_runtime/discovery/
mock.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use super::{
5    Discovery, DiscoveryEvent, DiscoveryInstance, DiscoveryInstanceId, DiscoveryQuery,
6    DiscoverySpec, DiscoveryStream, validate_event_source_reregistration,
7};
8use anyhow::Result;
9use async_trait::async_trait;
10use std::sync::{Arc, Mutex};
11use tokio_util::sync::CancellationToken;
12
13/// Shared in-memory registry for mock discovery
14#[derive(Clone, Default)]
15pub struct SharedMockRegistry {
16    instances: Arc<Mutex<Vec<DiscoveryInstance>>>,
17}
18
19impl SharedMockRegistry {
20    pub fn new() -> Self {
21        Self::default()
22    }
23}
24
25/// Mock implementation of Discovery for testing
26/// We can potentially remove this once we have KVStoreDiscovery fully tested
27pub struct MockDiscovery {
28    instance_id: u64,
29    registry: SharedMockRegistry,
30}
31
32impl MockDiscovery {
33    pub fn new(instance_id: Option<u64>, registry: SharedMockRegistry) -> Self {
34        let instance_id = instance_id.unwrap_or_else(|| {
35            use std::sync::atomic::{AtomicU64, Ordering};
36            static COUNTER: AtomicU64 = AtomicU64::new(1);
37            COUNTER.fetch_add(1, Ordering::SeqCst)
38        });
39
40        Self {
41            instance_id,
42            registry,
43        }
44    }
45}
46
47/// Helper function to check if an instance matches a discovery query
48fn matches_query(instance: &DiscoveryInstance, query: &DiscoveryQuery) -> bool {
49    match (instance, query) {
50        // Endpoint matching
51        (DiscoveryInstance::Endpoint(_), DiscoveryQuery::AllEndpoints) => true,
52        (DiscoveryInstance::Endpoint(inst), DiscoveryQuery::NamespacedEndpoints { namespace }) => {
53            &inst.namespace == namespace
54        }
55        (
56            DiscoveryInstance::Endpoint(inst),
57            DiscoveryQuery::ComponentEndpoints {
58                namespace,
59                component,
60            },
61        ) => &inst.namespace == namespace && &inst.component == component,
62        (
63            DiscoveryInstance::Endpoint(inst),
64            DiscoveryQuery::Endpoint {
65                namespace,
66                component,
67                endpoint,
68            },
69        ) => {
70            &inst.namespace == namespace
71                && &inst.component == component
72                && &inst.endpoint == endpoint
73        }
74
75        // Model matching
76        (DiscoveryInstance::Model { .. }, DiscoveryQuery::AllModels) => true,
77        (
78            DiscoveryInstance::Model {
79                namespace: inst_ns, ..
80            },
81            DiscoveryQuery::NamespacedModels { namespace },
82        ) => inst_ns == namespace,
83        (
84            DiscoveryInstance::Model {
85                namespace: inst_ns,
86                component: inst_comp,
87                ..
88            },
89            DiscoveryQuery::ComponentModels {
90                namespace,
91                component,
92            },
93        ) => inst_ns == namespace && inst_comp == component,
94        (
95            DiscoveryInstance::Model {
96                namespace: inst_ns,
97                component: inst_comp,
98                endpoint: inst_ep,
99                ..
100            },
101            DiscoveryQuery::EndpointModels {
102                namespace,
103                component,
104                endpoint,
105            },
106        ) => inst_ns == namespace && inst_comp == component && inst_ep == endpoint,
107
108        // EventChannel matching - unified query
109        (
110            DiscoveryInstance::EventChannel {
111                scope: inst_scope,
112                topic: inst_topic,
113                ..
114            },
115            DiscoveryQuery::EventChannels(query),
116        ) => {
117            query.scope.as_ref().is_none_or(|scope| scope == inst_scope)
118                && query.topic.as_ref().is_none_or(|t| t == inst_topic)
119        }
120
121        (
122            DiscoveryInstance::EventSource {
123                scope: inst_scope,
124                topic: inst_topic,
125                ..
126            },
127            DiscoveryQuery::EventSources(query),
128        ) => {
129            query.scope.as_ref().is_none_or(|scope| scope == inst_scope)
130                && query.topic.as_ref().is_none_or(|t| t == inst_topic)
131        }
132
133        // Cross-type matches return false
134        (
135            DiscoveryInstance::Endpoint(_),
136            DiscoveryQuery::AllModels
137            | DiscoveryQuery::NamespacedModels { .. }
138            | DiscoveryQuery::ComponentModels { .. }
139            | DiscoveryQuery::EndpointModels { .. }
140            | DiscoveryQuery::EventChannels(_)
141            | DiscoveryQuery::EventSources(_),
142        ) => false,
143        (
144            DiscoveryInstance::Model { .. },
145            DiscoveryQuery::AllEndpoints
146            | DiscoveryQuery::NamespacedEndpoints { .. }
147            | DiscoveryQuery::ComponentEndpoints { .. }
148            | DiscoveryQuery::Endpoint { .. }
149            | DiscoveryQuery::EventChannels(_)
150            | DiscoveryQuery::EventSources(_),
151        ) => false,
152        (
153            DiscoveryInstance::EventChannel { .. },
154            DiscoveryQuery::AllEndpoints
155            | DiscoveryQuery::NamespacedEndpoints { .. }
156            | DiscoveryQuery::ComponentEndpoints { .. }
157            | DiscoveryQuery::Endpoint { .. }
158            | DiscoveryQuery::AllModels
159            | DiscoveryQuery::NamespacedModels { .. }
160            | DiscoveryQuery::ComponentModels { .. }
161            | DiscoveryQuery::EndpointModels { .. },
162        ) => false,
163        (DiscoveryInstance::EventChannel { .. }, DiscoveryQuery::EventSources(_)) => false,
164        (
165            DiscoveryInstance::EventSource { .. },
166            DiscoveryQuery::AllEndpoints
167            | DiscoveryQuery::NamespacedEndpoints { .. }
168            | DiscoveryQuery::ComponentEndpoints { .. }
169            | DiscoveryQuery::Endpoint { .. }
170            | DiscoveryQuery::AllModels
171            | DiscoveryQuery::NamespacedModels { .. }
172            | DiscoveryQuery::ComponentModels { .. }
173            | DiscoveryQuery::EndpointModels { .. }
174            | DiscoveryQuery::EventChannels(_),
175        ) => false,
176    }
177}
178
179#[async_trait]
180impl Discovery for MockDiscovery {
181    fn instance_id(&self) -> u64 {
182        self.instance_id
183    }
184
185    async fn register_internal(&self, spec: DiscoverySpec) -> Result<DiscoveryInstance> {
186        let instance = spec.into_instance(self.instance_id);
187        let mut instances = self.registry.instances.lock().unwrap();
188        if matches!(&instance, DiscoveryInstance::EventSource { .. })
189            && let Some(existing) = instances
190                .iter()
191                .find(|existing| existing.id() == instance.id())
192        {
193            validate_event_source_reregistration(existing, &instance)?;
194            return Ok(existing.clone());
195        }
196        instances.push(instance.clone());
197
198        Ok(instance)
199    }
200
201    async fn unregister(&self, instance: DiscoveryInstance) -> Result<()> {
202        let target_id = instance.id();
203
204        self.registry
205            .instances
206            .lock()
207            .unwrap()
208            .retain(|i| i.id() != target_id);
209
210        Ok(())
211    }
212
213    async fn list(&self, query: DiscoveryQuery) -> Result<Vec<DiscoveryInstance>> {
214        let instances = self.registry.instances.lock().unwrap();
215        Ok(instances
216            .iter()
217            .filter(|instance| matches_query(instance, &query))
218            .cloned()
219            .collect())
220    }
221
222    async fn list_and_watch(
223        &self,
224        query: DiscoveryQuery,
225        _cancel_token: Option<CancellationToken>,
226    ) -> Result<DiscoveryStream> {
227        use std::collections::HashSet;
228
229        let registry = self.registry.clone();
230
231        let stream = async_stream::stream! {
232            let mut known_instances: HashSet<DiscoveryInstanceId> = HashSet::new();
233
234            loop {
235                let current: Vec<_> = {
236                    let instances = registry.instances.lock().unwrap();
237                    instances
238                        .iter()
239                        .filter(|instance| matches_query(instance, &query))
240                        .cloned()
241                        .collect()
242                };
243
244                let current_ids: HashSet<DiscoveryInstanceId> = current.iter().map(|i| i.id()).collect();
245
246                // Emit Added events for new instances
247                for instance in current {
248                    let id = instance.id();
249                    if known_instances.insert(id) {
250                        yield Ok(DiscoveryEvent::Added(instance));
251                    }
252                }
253
254                // Emit Removed events for instances that are gone
255                for id in known_instances.difference(&current_ids).cloned().collect::<Vec<_>>() {
256                    known_instances.remove(&id);
257                    yield Ok(DiscoveryEvent::Removed(id));
258                }
259
260                tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
261            }
262        };
263
264        Ok(Box::pin(stream))
265    }
266}
267
268#[cfg(test)]
269mod tests {
270    use super::*;
271    use futures::StreamExt;
272
273    fn model_spec(
274        namespace: &str,
275        component: &str,
276        endpoint: &str,
277        model_name: &str,
278    ) -> DiscoverySpec {
279        DiscoverySpec::Model {
280            namespace: namespace.to_string(),
281            component: component.to_string(),
282            endpoint: endpoint.to_string(),
283            card_json: serde_json::json!({
284                "display_name": model_name,
285            }),
286            model_suffix: None,
287        }
288    }
289
290    fn lora_model_spec(
291        namespace: &str,
292        component: &str,
293        endpoint: &str,
294        model_name: &str,
295        source_path: &str,
296        lora_name: &str,
297    ) -> DiscoverySpec {
298        DiscoverySpec::Model {
299            namespace: namespace.to_string(),
300            component: component.to_string(),
301            endpoint: endpoint.to_string(),
302            card_json: serde_json::json!({
303                "display_name": model_name,
304                "source_path": source_path,
305                "lora": {
306                    "name": lora_name,
307                },
308            }),
309            model_suffix: Some(lora_name.to_string()),
310        }
311    }
312
313    #[tokio::test]
314    async fn test_mock_discovery_add_and_remove() {
315        let registry = SharedMockRegistry::new();
316        let client1 = MockDiscovery::new(Some(1), registry.clone());
317        let client2 = MockDiscovery::new(Some(2), registry.clone());
318
319        let spec = DiscoverySpec::Endpoint {
320            namespace: "test-ns".to_string(),
321            component: "test-comp".to_string(),
322            endpoint: "test-ep".to_string(),
323            transport: crate::component::TransportType::Nats("test-subject".to_string()),
324            device_type: None,
325            request_plane_codec: None,
326        };
327
328        let query = DiscoveryQuery::Endpoint {
329            namespace: "test-ns".to_string(),
330            component: "test-comp".to_string(),
331            endpoint: "test-ep".to_string(),
332        };
333
334        // Start watching
335        let mut stream = client1.list_and_watch(query.clone(), None).await.unwrap();
336
337        // Add first instance
338        let instance1 = client1.register(spec.clone()).await.unwrap();
339
340        let event = stream.next().await.unwrap().unwrap();
341        match event {
342            DiscoveryEvent::Added(DiscoveryInstance::Endpoint(inst)) => {
343                assert_eq!(inst.instance_id, 1);
344            }
345            _ => panic!("Expected Added event for instance-1"),
346        }
347
348        // Add second instance
349        client2.register(spec.clone()).await.unwrap();
350
351        let event = stream.next().await.unwrap().unwrap();
352        match event {
353            DiscoveryEvent::Added(DiscoveryInstance::Endpoint(inst)) => {
354                assert_eq!(inst.instance_id, 2);
355            }
356            _ => panic!("Expected Added event for instance-2"),
357        }
358
359        // Remove first instance
360        client1.unregister(instance1).await.unwrap();
361
362        let event = stream.next().await.unwrap().unwrap();
363        match event {
364            DiscoveryEvent::Removed(id) => {
365                let endpoint_id = id.extract_endpoint_id().expect("Expected endpoint removal");
366                assert_eq!(endpoint_id.instance_id, 1);
367            }
368            _ => panic!("Expected Removed event for instance-1"),
369        }
370    }
371
372    #[tokio::test]
373    async fn event_source_removal_is_publisher_specific() {
374        use crate::discovery::{EventScope, EventSourceQuery};
375
376        let client = MockDiscovery::new(Some(42), SharedMockRegistry::new());
377        let endpoint = crate::protocols::EndpointId {
378            namespace: "workers".to_string(),
379            component: "backend".to_string(),
380            name: "kv-state".to_string(),
381        };
382        let query = DiscoveryQuery::EventSources(EventSourceQuery::endpoint_topic(
383            endpoint.clone(),
384            "kv-events",
385        ));
386        let spec = |publisher_id, worker_id| DiscoverySpec::EventSource {
387            scope: EventScope::Endpoint {
388                endpoint: endpoint.clone(),
389            },
390            topic: "kv-events".to_string(),
391            publisher_id,
392            metadata: serde_json::json!({"worker_id": worker_id, "dp_rank": 0}),
393        };
394
395        let old = client.register(spec(100, 7)).await.unwrap();
396        assert_eq!(client.register(spec(100, 7)).await.unwrap(), old);
397        assert!(client.register(spec(100, 8)).await.is_err());
398        assert_eq!(client.list(query.clone()).await.unwrap(), vec![old.clone()]);
399
400        let current = client.register(spec(205, 7)).await.unwrap();
401        assert_eq!(client.list(query.clone()).await.unwrap().len(), 2);
402
403        client.unregister(old).await.unwrap();
404        assert_eq!(client.list(query).await.unwrap(), vec![current]);
405    }
406
407    #[tokio::test]
408    async fn register_allows_same_model_name_on_same_endpoint() {
409        let registry = SharedMockRegistry::new();
410        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
411        let discovery2 = MockDiscovery::new(Some(2), registry);
412        let spec = model_spec("ns", "comp", "generate", "model-a");
413
414        discovery1.register(spec.clone()).await.unwrap();
415        discovery2.register(spec).await.unwrap();
416
417        let instances = discovery1
418            .list(DiscoveryQuery::EndpointModels {
419                namespace: "ns".to_string(),
420                component: "comp".to_string(),
421                endpoint: "generate".to_string(),
422            })
423            .await
424            .unwrap();
425        assert_eq!(instances.len(), 2);
426    }
427
428    #[tokio::test]
429    async fn register_rejects_distinct_base_cards_with_same_source_path_on_same_endpoint() {
430        let registry = SharedMockRegistry::new();
431        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
432        let discovery2 = MockDiscovery::new(Some(2), registry);
433        let spec = |display_name: &str| DiscoverySpec::Model {
434            namespace: "ns".to_string(),
435            component: "comp".to_string(),
436            endpoint: "generate".to_string(),
437            card_json: serde_json::json!({
438                "display_name": display_name,
439                "source_path": "org/base-model",
440            }),
441            model_suffix: None,
442        };
443
444        discovery1.register(spec("public-name-a")).await.unwrap();
445        let err = discovery2
446            .register(spec("public-name-b"))
447            .await
448            .unwrap_err();
449
450        assert!(err.to_string().contains(
451            "Cannot register model 'public-name-b' on endpoint 'ns/comp/generate': a different model 'public-name-a' is already registered there"
452        ));
453
454        let instances = discovery1
455            .list(DiscoveryQuery::EndpointModels {
456                namespace: "ns".to_string(),
457                component: "comp".to_string(),
458                endpoint: "generate".to_string(),
459            })
460            .await
461            .unwrap();
462        assert_eq!(instances.len(), 1);
463    }
464
465    #[tokio::test]
466    async fn register_rejects_different_model_name_on_same_endpoint() {
467        let registry = SharedMockRegistry::new();
468        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
469        let discovery2 = MockDiscovery::new(Some(2), registry);
470
471        discovery1
472            .register(model_spec("ns", "comp", "generate", "model-a"))
473            .await
474            .unwrap();
475
476        let err = discovery2
477            .register(model_spec("ns", "comp", "generate", "model-b"))
478            .await
479            .unwrap_err();
480
481        assert!(err.to_string().contains(
482            "Cannot register model 'model-b' on endpoint 'ns/comp/generate': a different model 'model-a' is already registered there"
483        ));
484
485        let instances = discovery1
486            .list(DiscoveryQuery::EndpointModels {
487                namespace: "ns".to_string(),
488                component: "comp".to_string(),
489                endpoint: "generate".to_string(),
490            })
491            .await
492            .unwrap();
493        assert_eq!(instances.len(), 1);
494    }
495
496    #[tokio::test]
497    async fn register_allows_different_model_names_on_different_endpoints() {
498        let registry = SharedMockRegistry::new();
499        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
500        let discovery2 = MockDiscovery::new(Some(2), registry);
501
502        discovery1
503            .register(model_spec("ns", "comp", "generate-a", "model-a"))
504            .await
505            .unwrap();
506        discovery2
507            .register(model_spec("ns", "comp", "generate-b", "model-b"))
508            .await
509            .unwrap();
510    }
511
512    #[tokio::test]
513    async fn register_allows_lora_adapter_on_same_endpoint() {
514        let registry = SharedMockRegistry::new();
515        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
516        let discovery2 = MockDiscovery::new(Some(2), registry);
517
518        discovery1
519            .register(DiscoverySpec::Model {
520                namespace: "ns".to_string(),
521                component: "comp".to_string(),
522                endpoint: "generate".to_string(),
523                card_json: serde_json::json!({
524                    "display_name": "base-model",
525                    "source_path": "base-repo",
526                }),
527                model_suffix: None,
528            })
529            .await
530            .unwrap();
531
532        discovery2
533            .register(lora_model_spec(
534                "ns",
535                "comp",
536                "generate",
537                "adapter-a",
538                "base-repo",
539                "adapter-a",
540            ))
541            .await
542            .unwrap();
543    }
544
545    #[tokio::test]
546    async fn register_rejects_lora_adapter_for_different_base_model() {
547        let registry = SharedMockRegistry::new();
548        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
549        let discovery2 = MockDiscovery::new(Some(2), registry);
550
551        discovery1
552            .register(DiscoverySpec::Model {
553                namespace: "ns".to_string(),
554                component: "comp".to_string(),
555                endpoint: "generate".to_string(),
556                card_json: serde_json::json!({
557                    "display_name": "base-model",
558                    "source_path": "base-repo",
559                }),
560                model_suffix: None,
561            })
562            .await
563            .unwrap();
564
565        let err = discovery2
566            .register(lora_model_spec(
567                "ns",
568                "comp",
569                "generate",
570                "adapter-a",
571                "other-base-repo",
572                "adapter-a",
573            ))
574            .await
575            .unwrap_err();
576
577        assert!(err.to_string().contains(
578            "Cannot register model 'adapter-a' on endpoint 'ns/comp/generate': a different model 'base-model' is already registered there"
579        ));
580    }
581}