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, ModelCardInstanceId, model_with_updated_taints,
7    reconcile_discovery_snapshot, validate_event_source_reregistration,
8    validate_model_reregistration,
9};
10use anyhow::Result;
11use async_trait::async_trait;
12use std::collections::{HashMap, HashSet};
13use std::sync::{Arc, Mutex};
14use tokio_util::sync::CancellationToken;
15
16/// Shared in-memory registry for mock discovery
17#[derive(Clone, Default)]
18pub struct SharedMockRegistry {
19    instances: Arc<Mutex<Vec<DiscoveryInstance>>>,
20}
21
22impl SharedMockRegistry {
23    pub fn new() -> Self {
24        Self::default()
25    }
26}
27
28/// Mock implementation of Discovery for testing
29/// We can potentially remove this once we have KVStoreDiscovery fully tested
30pub struct MockDiscovery {
31    instance_id: u64,
32    registry: SharedMockRegistry,
33}
34
35impl MockDiscovery {
36    pub fn new(instance_id: Option<u64>, registry: SharedMockRegistry) -> Self {
37        let instance_id = instance_id.unwrap_or_else(|| {
38            use std::sync::atomic::{AtomicU64, Ordering};
39            static COUNTER: AtomicU64 = AtomicU64::new(1);
40            COUNTER.fetch_add(1, Ordering::SeqCst)
41        });
42
43        Self {
44            instance_id,
45            registry,
46        }
47    }
48}
49
50/// Helper function to check if an instance matches a discovery query
51fn matches_query(instance: &DiscoveryInstance, query: &DiscoveryQuery) -> bool {
52    match (instance, query) {
53        // Endpoint matching
54        (DiscoveryInstance::Endpoint(_), DiscoveryQuery::AllEndpoints) => true,
55        (DiscoveryInstance::Endpoint(inst), DiscoveryQuery::NamespacedEndpoints { namespace }) => {
56            &inst.namespace == namespace
57        }
58        (
59            DiscoveryInstance::Endpoint(inst),
60            DiscoveryQuery::ComponentEndpoints {
61                namespace,
62                component,
63            },
64        ) => &inst.namespace == namespace && &inst.component == component,
65        (
66            DiscoveryInstance::Endpoint(inst),
67            DiscoveryQuery::Endpoint {
68                namespace,
69                component,
70                endpoint,
71            },
72        ) => {
73            &inst.namespace == namespace
74                && &inst.component == component
75                && &inst.endpoint == endpoint
76        }
77
78        // Model matching
79        (DiscoveryInstance::Model { .. }, DiscoveryQuery::AllModels) => true,
80        (
81            DiscoveryInstance::Model {
82                namespace: inst_ns, ..
83            },
84            DiscoveryQuery::NamespacedModels { namespace },
85        ) => inst_ns == namespace,
86        (
87            DiscoveryInstance::Model {
88                namespace: inst_ns,
89                component: inst_comp,
90                ..
91            },
92            DiscoveryQuery::ComponentModels {
93                namespace,
94                component,
95            },
96        ) => inst_ns == namespace && inst_comp == component,
97        (
98            DiscoveryInstance::Model {
99                namespace: inst_ns,
100                component: inst_comp,
101                endpoint: inst_ep,
102                ..
103            },
104            DiscoveryQuery::EndpointModels {
105                namespace,
106                component,
107                endpoint,
108            },
109        ) => inst_ns == namespace && inst_comp == component && inst_ep == endpoint,
110
111        // EventChannel matching - unified query
112        (
113            DiscoveryInstance::EventChannel {
114                scope: inst_scope,
115                topic: inst_topic,
116                ..
117            },
118            DiscoveryQuery::EventChannels(query),
119        ) => {
120            query.scope.as_ref().is_none_or(|scope| scope == inst_scope)
121                && query.topic.as_ref().is_none_or(|t| t == inst_topic)
122        }
123
124        (
125            DiscoveryInstance::EventSource {
126                scope: inst_scope,
127                topic: inst_topic,
128                ..
129            },
130            DiscoveryQuery::EventSources(query),
131        ) => {
132            query.scope.as_ref().is_none_or(|scope| scope == inst_scope)
133                && query.topic.as_ref().is_none_or(|t| t == inst_topic)
134        }
135
136        // Cross-type matches return false
137        (
138            DiscoveryInstance::Endpoint(_),
139            DiscoveryQuery::AllModels
140            | DiscoveryQuery::NamespacedModels { .. }
141            | DiscoveryQuery::ComponentModels { .. }
142            | DiscoveryQuery::EndpointModels { .. }
143            | DiscoveryQuery::EventChannels(_)
144            | DiscoveryQuery::EventSources(_),
145        ) => false,
146        (
147            DiscoveryInstance::Model { .. },
148            DiscoveryQuery::AllEndpoints
149            | DiscoveryQuery::NamespacedEndpoints { .. }
150            | DiscoveryQuery::ComponentEndpoints { .. }
151            | DiscoveryQuery::Endpoint { .. }
152            | DiscoveryQuery::EventChannels(_)
153            | DiscoveryQuery::EventSources(_),
154        ) => false,
155        (
156            DiscoveryInstance::EventChannel { .. },
157            DiscoveryQuery::AllEndpoints
158            | DiscoveryQuery::NamespacedEndpoints { .. }
159            | DiscoveryQuery::ComponentEndpoints { .. }
160            | DiscoveryQuery::Endpoint { .. }
161            | DiscoveryQuery::AllModels
162            | DiscoveryQuery::NamespacedModels { .. }
163            | DiscoveryQuery::ComponentModels { .. }
164            | DiscoveryQuery::EndpointModels { .. },
165        ) => false,
166        (DiscoveryInstance::EventChannel { .. }, DiscoveryQuery::EventSources(_)) => false,
167        (
168            DiscoveryInstance::EventSource { .. },
169            DiscoveryQuery::AllEndpoints
170            | DiscoveryQuery::NamespacedEndpoints { .. }
171            | DiscoveryQuery::ComponentEndpoints { .. }
172            | DiscoveryQuery::Endpoint { .. }
173            | DiscoveryQuery::AllModels
174            | DiscoveryQuery::NamespacedModels { .. }
175            | DiscoveryQuery::ComponentModels { .. }
176            | DiscoveryQuery::EndpointModels { .. }
177            | DiscoveryQuery::EventChannels(_),
178        ) => false,
179    }
180}
181
182#[async_trait]
183impl Discovery for MockDiscovery {
184    fn instance_id(&self) -> u64 {
185        self.instance_id
186    }
187
188    async fn register_internal(&self, spec: DiscoverySpec) -> Result<DiscoveryInstance> {
189        let instance = spec.into_instance(self.instance_id);
190        let instance_id = instance.id();
191        let mut instances = self.registry.instances.lock().unwrap();
192        if let Some(existing) = instances
193            .iter_mut()
194            .find(|existing| existing.id() == instance_id)
195        {
196            match &instance {
197                DiscoveryInstance::Endpoint(_) => {
198                    *existing = instance.clone();
199                    return Ok(instance);
200                }
201                DiscoveryInstance::EventSource { .. } => {
202                    validate_event_source_reregistration(existing, &instance)?;
203                    return Ok(existing.clone());
204                }
205                DiscoveryInstance::Model { .. } => {
206                    validate_model_reregistration(existing, &instance)?;
207                    return Ok(existing.clone());
208                }
209                DiscoveryInstance::EventChannel { .. } => {}
210            }
211        }
212        instances.push(instance.clone());
213
214        Ok(instance)
215    }
216
217    async fn update_model_taints_internal(
218        &self,
219        id: ModelCardInstanceId,
220        taints: HashSet<String>,
221    ) -> Result<()> {
222        let target_id = DiscoveryInstanceId::Model(id);
223        let mut instances = self.registry.instances.lock().unwrap();
224        let existing = instances
225            .iter_mut()
226            .find(|existing| existing.id() == target_id)
227            .ok_or_else(|| {
228                anyhow::anyhow!("model discovery record {target_id:?} is not registered")
229            })?;
230        *existing = model_with_updated_taints(existing, taints)?;
231        Ok(())
232    }
233
234    async fn unregister(&self, instance: DiscoveryInstance) -> Result<()> {
235        let target_id = instance.id();
236
237        self.registry
238            .instances
239            .lock()
240            .unwrap()
241            .retain(|i| i.id() != target_id);
242
243        Ok(())
244    }
245
246    async fn list(&self, query: DiscoveryQuery) -> Result<Vec<DiscoveryInstance>> {
247        let instances = self.registry.instances.lock().unwrap();
248        Ok(instances
249            .iter()
250            .filter(|instance| matches_query(instance, &query))
251            .cloned()
252            .collect())
253    }
254
255    async fn list_and_watch(
256        &self,
257        query: DiscoveryQuery,
258        _cancel_token: Option<CancellationToken>,
259    ) -> Result<DiscoveryStream> {
260        let registry = self.registry.clone();
261
262        let stream = async_stream::stream! {
263            let mut known_instances = HashMap::<DiscoveryInstanceId, DiscoveryInstance>::new();
264
265            loop {
266                let current: HashMap<DiscoveryInstanceId, DiscoveryInstance> = {
267                    let instances = registry.instances.lock().unwrap();
268                    instances
269                        .iter()
270                        .filter(|instance| matches_query(instance, &query))
271                        .cloned()
272                        .map(|instance| (instance.id(), instance))
273                        .collect()
274                };
275
276                let (events, reconciled) =
277                    reconcile_discovery_snapshot(&known_instances, current);
278                for event in events {
279                    yield Ok(event);
280                }
281
282                known_instances = reconciled;
283                tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
284            }
285        };
286
287        Ok(Box::pin(stream))
288    }
289}
290
291#[cfg(test)]
292mod tests {
293    use super::*;
294    use crate::component::TransportType;
295    use futures::StreamExt;
296    use tokio::time::{Duration, timeout};
297
298    #[tokio::test]
299    async fn watch_emits_same_id_endpoint_update() {
300        let client = MockDiscovery::new(Some(1), SharedMockRegistry::new());
301        let query = DiscoveryQuery::Endpoint {
302            namespace: "ns".to_string(),
303            component: "component".to_string(),
304            endpoint: "endpoint".to_string(),
305        };
306        let spec = |transport: &str| DiscoverySpec::Endpoint {
307            namespace: "ns".to_string(),
308            component: "component".to_string(),
309            endpoint: "endpoint".to_string(),
310            transport: TransportType::Tcp(transport.to_string()),
311            device_type: None,
312            request_plane_codec: None,
313        };
314        let mut stream = client.list_and_watch(query.clone(), None).await.unwrap();
315
316        let original = client.register(spec("127.0.0.1:8000")).await.unwrap();
317        let event = timeout(Duration::from_secs(1), stream.next())
318            .await
319            .expect("mock watch should emit the initial instance")
320            .unwrap()
321            .unwrap();
322        assert_eq!(event, DiscoveryEvent::Added(original));
323
324        let updated = client.register(spec("127.0.0.1:9000")).await.unwrap();
325        let event = timeout(Duration::from_secs(1), stream.next())
326            .await
327            .expect("mock watch should emit the updated instance")
328            .unwrap()
329            .unwrap();
330
331        assert_eq!(event, DiscoveryEvent::Added(updated.clone()));
332        assert_eq!(client.list(query).await.unwrap(), vec![updated]);
333    }
334
335    fn model_spec(
336        namespace: &str,
337        component: &str,
338        endpoint: &str,
339        model_name: &str,
340    ) -> DiscoverySpec {
341        DiscoverySpec::Model {
342            namespace: namespace.to_string(),
343            component: component.to_string(),
344            endpoint: endpoint.to_string(),
345            card_json: serde_json::json!({
346                "display_name": model_name,
347            }),
348            model_suffix: None,
349        }
350    }
351
352    fn lora_model_spec(
353        namespace: &str,
354        component: &str,
355        endpoint: &str,
356        model_name: &str,
357        source_path: &str,
358        lora_name: &str,
359    ) -> DiscoverySpec {
360        DiscoverySpec::Model {
361            namespace: namespace.to_string(),
362            component: component.to_string(),
363            endpoint: endpoint.to_string(),
364            card_json: serde_json::json!({
365                "display_name": model_name,
366                "source_path": source_path,
367                "lora": {
368                    "name": lora_name,
369                },
370            }),
371            model_suffix: Some(lora_name.to_string()),
372        }
373    }
374
375    #[tokio::test]
376    async fn model_taint_updates_use_the_authoritative_registry() {
377        let client = MockDiscovery::new(Some(7), SharedMockRegistry::new());
378        let model = client
379            .register(DiscoverySpec::Model {
380                namespace: "ns".to_string(),
381                component: "worker".to_string(),
382                endpoint: "generate".to_string(),
383                card_json: serde_json::json!({
384                    "display_name": "model",
385                    "runtime_config": {"taints": ["a"]}
386                }),
387                model_suffix: None,
388            })
389            .await
390            .unwrap();
391        let DiscoveryInstanceId::Model(id) = model.id() else {
392            unreachable!()
393        };
394
395        client
396            .update_model_taints(id.clone(), HashSet::from(["b".to_string()]))
397            .await
398            .unwrap();
399        client
400            .update_model_taints(id.clone(), HashSet::from(["a".to_string()]))
401            .await
402            .unwrap();
403
404        let stored = client
405            .list(DiscoveryQuery::EndpointModels {
406                namespace: "ns".to_string(),
407                component: "worker".to_string(),
408                endpoint: "generate".to_string(),
409            })
410            .await
411            .unwrap()
412            .pop()
413            .unwrap();
414        let DiscoveryInstance::Model { card_json, .. } = stored else {
415            unreachable!()
416        };
417        assert_eq!(
418            card_json["runtime_config"]["taints"],
419            serde_json::json!(["a"])
420        );
421
422        client.unregister(model).await.unwrap();
423        assert!(
424            client
425                .update_model_taints(id, HashSet::new())
426                .await
427                .is_err()
428        );
429    }
430
431    #[tokio::test]
432    async fn same_id_model_registration_preserves_updated_taints_without_duplicates() {
433        let client = MockDiscovery::new(Some(7), SharedMockRegistry::new());
434        let spec = DiscoverySpec::Model {
435            namespace: "ns".to_string(),
436            component: "worker".to_string(),
437            endpoint: "generate".to_string(),
438            card_json: serde_json::json!({
439                "display_name": "model",
440                "runtime_config": {"taints": ["initial"]}
441            }),
442            model_suffix: None,
443        };
444        let original = client.register(spec.clone()).await.unwrap();
445        let DiscoveryInstanceId::Model(id) = original.id() else {
446            unreachable!()
447        };
448        client
449            .update_model_taints(id, HashSet::from(["updated".to_string()]))
450            .await
451            .unwrap();
452
453        let replayed = client.register(spec).await.unwrap();
454        let models = client
455            .list(DiscoveryQuery::EndpointModels {
456                namespace: "ns".to_string(),
457                component: "worker".to_string(),
458                endpoint: "generate".to_string(),
459            })
460            .await
461            .unwrap();
462
463        assert_eq!(models, vec![replayed.clone()]);
464        let DiscoveryInstance::Model { card_json, .. } = replayed else {
465            unreachable!()
466        };
467        assert_eq!(
468            card_json["runtime_config"]["taints"],
469            serde_json::json!(["updated"])
470        );
471    }
472
473    #[tokio::test]
474    async fn model_taint_update_rejects_foreign_worker_id() {
475        let client = MockDiscovery::new(Some(7), SharedMockRegistry::new());
476        let foreign_id = ModelCardInstanceId {
477            namespace: "ns".to_string(),
478            component: "worker".to_string(),
479            endpoint: "generate".to_string(),
480            instance_id: 8,
481            model_suffix: None,
482        };
483
484        let error = client
485            .update_model_taints(foreign_id, HashSet::new())
486            .await
487            .unwrap_err();
488
489        assert!(
490            error
491                .to_string()
492                .contains("this discovery client owns worker 7")
493        );
494    }
495
496    #[tokio::test]
497    async fn test_mock_discovery_add_and_remove() {
498        let registry = SharedMockRegistry::new();
499        let client1 = MockDiscovery::new(Some(1), registry.clone());
500        let client2 = MockDiscovery::new(Some(2), registry.clone());
501
502        let spec = DiscoverySpec::Endpoint {
503            namespace: "test-ns".to_string(),
504            component: "test-comp".to_string(),
505            endpoint: "test-ep".to_string(),
506            transport: crate::component::TransportType::Nats("test-subject".to_string()),
507            device_type: None,
508            request_plane_codec: None,
509        };
510
511        let query = DiscoveryQuery::Endpoint {
512            namespace: "test-ns".to_string(),
513            component: "test-comp".to_string(),
514            endpoint: "test-ep".to_string(),
515        };
516
517        // Start watching
518        let mut stream = client1.list_and_watch(query.clone(), None).await.unwrap();
519
520        // Add first instance
521        let instance1 = client1.register(spec.clone()).await.unwrap();
522
523        let event = stream.next().await.unwrap().unwrap();
524        match event {
525            DiscoveryEvent::Added(DiscoveryInstance::Endpoint(inst)) => {
526                assert_eq!(inst.instance_id, 1);
527            }
528            _ => panic!("Expected Added event for instance-1"),
529        }
530
531        // Add second instance
532        client2.register(spec.clone()).await.unwrap();
533
534        let event = stream.next().await.unwrap().unwrap();
535        match event {
536            DiscoveryEvent::Added(DiscoveryInstance::Endpoint(inst)) => {
537                assert_eq!(inst.instance_id, 2);
538            }
539            _ => panic!("Expected Added event for instance-2"),
540        }
541
542        // Remove first instance
543        client1.unregister(instance1).await.unwrap();
544
545        let event = stream.next().await.unwrap().unwrap();
546        match event {
547            DiscoveryEvent::Removed(id) => {
548                let endpoint_id = id.extract_endpoint_id().expect("Expected endpoint removal");
549                assert_eq!(endpoint_id.instance_id, 1);
550            }
551            _ => panic!("Expected Removed event for instance-1"),
552        }
553    }
554
555    #[tokio::test]
556    async fn event_source_removal_is_publisher_specific() {
557        use crate::discovery::{EventScope, EventSourceQuery};
558
559        let client = MockDiscovery::new(Some(42), SharedMockRegistry::new());
560        let endpoint = crate::protocols::EndpointId {
561            namespace: "workers".to_string(),
562            component: "backend".to_string(),
563            name: "kv-state".to_string(),
564        };
565        let query = DiscoveryQuery::EventSources(EventSourceQuery::endpoint_topic(
566            endpoint.clone(),
567            "kv-events",
568        ));
569        let spec = |publisher_id, worker_id| DiscoverySpec::EventSource {
570            scope: EventScope::Endpoint {
571                endpoint: endpoint.clone(),
572            },
573            topic: "kv-events".to_string(),
574            publisher_id,
575            metadata: serde_json::json!({"worker_id": worker_id, "dp_rank": 0}),
576        };
577
578        let old = client.register(spec(100, 7)).await.unwrap();
579        assert_eq!(client.register(spec(100, 7)).await.unwrap(), old);
580        assert!(client.register(spec(100, 8)).await.is_err());
581        assert_eq!(client.list(query.clone()).await.unwrap(), vec![old.clone()]);
582
583        let current = client.register(spec(205, 7)).await.unwrap();
584        assert_eq!(client.list(query.clone()).await.unwrap().len(), 2);
585
586        client.unregister(old).await.unwrap();
587        assert_eq!(client.list(query).await.unwrap(), vec![current]);
588    }
589
590    #[tokio::test]
591    async fn register_allows_same_model_name_on_same_endpoint() {
592        let registry = SharedMockRegistry::new();
593        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
594        let discovery2 = MockDiscovery::new(Some(2), registry);
595        let spec = model_spec("ns", "comp", "generate", "model-a");
596
597        discovery1.register(spec.clone()).await.unwrap();
598        discovery2.register(spec).await.unwrap();
599
600        let instances = discovery1
601            .list(DiscoveryQuery::EndpointModels {
602                namespace: "ns".to_string(),
603                component: "comp".to_string(),
604                endpoint: "generate".to_string(),
605            })
606            .await
607            .unwrap();
608        assert_eq!(instances.len(), 2);
609    }
610
611    #[tokio::test]
612    async fn register_non_lora_alias_compatibility() {
613        for (name, source_a, source_b, compatible) in [
614            ("alias-b", Some("org/base"), Some("org/base"), true),
615            ("alias-a", Some("/mount/a"), Some("/mount/b"), true),
616            ("alias-b", Some("org/base"), Some("org/other"), false),
617            ("alias-b", Some("org/base"), None, false),
618            ("alias-b", None, Some("org/base"), false),
619            ("alias-b", Some(""), Some(""), false),
620        ] {
621            let registry = SharedMockRegistry::new();
622            let discovery1 = MockDiscovery::new(Some(1), registry.clone());
623            let discovery2 = MockDiscovery::new(Some(2), registry);
624            let spec = |display_name: &str, source_path: Option<&str>| DiscoverySpec::Model {
625                namespace: "ns".to_string(),
626                component: "comp".to_string(),
627                endpoint: "generate".to_string(),
628                card_json: serde_json::json!({
629                    "display_name": display_name,
630                    "source_path": source_path,
631                }),
632                model_suffix: None,
633            };
634            discovery1
635                .register(spec("alias-a", source_a))
636                .await
637                .unwrap();
638            let result = discovery2.register(spec(name, source_b)).await;
639            assert_eq!(
640                result.is_ok(),
641                compatible,
642                "{name}: {source_a:?}, {source_b:?}: {result:?}"
643            );
644            if let Err(err) = result {
645                assert!(
646                    err.to_string()
647                        .contains("a different model 'alias-a' is already registered there")
648                );
649            }
650            let instances = discovery1
651                .list(DiscoveryQuery::EndpointModels {
652                    namespace: "ns".to_string(),
653                    component: "comp".to_string(),
654                    endpoint: "generate".to_string(),
655                })
656                .await
657                .unwrap();
658            assert_eq!(instances.len(), if compatible { 2 } else { 1 });
659        }
660    }
661
662    #[tokio::test]
663    async fn register_shared_source_requires_disjoint_served_names() {
664        for (name_b, aliases_a, aliases_b, compatible) in [
665            ("b", vec!["b"], vec![], false),
666            ("b", vec![], vec!["a"], false),
667            ("b", vec!["shared"], vec!["shared"], false),
668            ("b", vec!["a", "extra-a"], vec!["b", "extra-b"], true),
669            ("a", vec!["shared"], vec!["shared"], true),
670        ] {
671            let registry = SharedMockRegistry::new();
672            let first = MockDiscovery::new(Some(1), registry.clone());
673            let second = MockDiscovery::new(Some(2), registry);
674            let spec = |name: &str, aliases: Vec<&str>| DiscoverySpec::Model {
675                namespace: "ns".into(),
676                component: "comp".into(),
677                endpoint: "generate".into(),
678                card_json: serde_json::json!({
679                    "display_name": name,
680                    "aliases": aliases,
681                    "source_path": "org/base",
682                }),
683                model_suffix: None,
684            };
685            let incumbent = first.register(spec("a", aliases_a)).await.unwrap();
686            let result = second.register(spec(name_b, aliases_b)).await;
687            assert_eq!(result.is_ok(), compatible, "{result:?}");
688            let instances = first
689                .list(DiscoveryQuery::EndpointModels {
690                    namespace: "ns".into(),
691                    component: "comp".into(),
692                    endpoint: "generate".into(),
693                })
694                .await
695                .unwrap();
696            assert!(instances.contains(&incumbent));
697            assert_eq!(instances.len(), if compatible { 2 } else { 1 });
698        }
699    }
700
701    #[tokio::test]
702    async fn register_checks_every_existing_model() {
703        // B is compatible with A by name and C by source, but A and C conflict.
704        for order in [[0, 1, 2], [2, 1, 0]] {
705            let registry = SharedMockRegistry::new();
706            let cards = [("a", "/mount/a"), ("a", "/mount/b"), ("b", "/mount/b")];
707            let mut accepted = Vec::new();
708            for (position, index) in order.into_iter().enumerate() {
709                let discovery = MockDiscovery::new(Some(index as u64 + 1), registry.clone());
710                let (name, source) = cards[index];
711                let result = discovery
712                    .register(DiscoverySpec::Model {
713                        namespace: "ns".into(),
714                        component: "comp".into(),
715                        endpoint: "generate".into(),
716                        card_json: serde_json::json!({"display_name": name, "source_path": source}),
717                        model_suffix: None,
718                    })
719                    .await;
720                if position < 2 {
721                    accepted.push(result.unwrap());
722                } else {
723                    assert!(result.is_err());
724                    let instances = discovery
725                        .list(DiscoveryQuery::EndpointModels {
726                            namespace: "ns".into(),
727                            component: "comp".into(),
728                            endpoint: "generate".into(),
729                        })
730                        .await
731                        .unwrap();
732                    assert_eq!(instances.len(), accepted.len());
733                    assert!(accepted.iter().all(|instance| instances.contains(instance)));
734                }
735            }
736        }
737    }
738
739    #[tokio::test]
740    async fn register_rejects_different_model_name_on_same_endpoint() {
741        let registry = SharedMockRegistry::new();
742        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
743        let discovery2 = MockDiscovery::new(Some(2), registry);
744
745        discovery1
746            .register(model_spec("ns", "comp", "generate", "model-a"))
747            .await
748            .unwrap();
749
750        let err = discovery2
751            .register(model_spec("ns", "comp", "generate", "model-b"))
752            .await
753            .unwrap_err();
754
755        assert!(err.to_string().contains(
756            "Cannot register model 'model-b' on endpoint 'ns/comp/generate': a different model 'model-a' is already registered there"
757        ));
758
759        let instances = discovery1
760            .list(DiscoveryQuery::EndpointModels {
761                namespace: "ns".to_string(),
762                component: "comp".to_string(),
763                endpoint: "generate".to_string(),
764            })
765            .await
766            .unwrap();
767        assert_eq!(instances.len(), 1);
768    }
769
770    #[tokio::test]
771    async fn register_allows_different_model_names_on_different_endpoints() {
772        let registry = SharedMockRegistry::new();
773        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
774        let discovery2 = MockDiscovery::new(Some(2), registry);
775
776        discovery1
777            .register(model_spec("ns", "comp", "generate-a", "model-a"))
778            .await
779            .unwrap();
780        discovery2
781            .register(model_spec("ns", "comp", "generate-b", "model-b"))
782            .await
783            .unwrap();
784    }
785
786    #[tokio::test]
787    async fn register_allows_lora_adapter_on_same_endpoint() {
788        let registry = SharedMockRegistry::new();
789        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
790        let discovery2 = MockDiscovery::new(Some(2), registry);
791
792        discovery1
793            .register(DiscoverySpec::Model {
794                namespace: "ns".to_string(),
795                component: "comp".to_string(),
796                endpoint: "generate".to_string(),
797                card_json: serde_json::json!({
798                    "display_name": "base-model",
799                    "source_path": "base-repo",
800                }),
801                model_suffix: None,
802            })
803            .await
804            .unwrap();
805
806        discovery2
807            .register(lora_model_spec(
808                "ns",
809                "comp",
810                "generate",
811                "adapter-a",
812                "base-repo",
813                "adapter-a",
814            ))
815            .await
816            .unwrap();
817    }
818
819    #[tokio::test]
820    async fn register_base_and_lora_require_distinct_served_names() {
821        for (adapter_name, compatible) in [("base", false), ("alias", false), ("adapter", true)] {
822            for adapter_first in [false, true] {
823                let registry = SharedMockRegistry::new();
824                let first = MockDiscovery::new(Some(1), registry.clone());
825                let second = MockDiscovery::new(Some(2), registry);
826                let base = DiscoverySpec::Model {
827                    namespace: "ns".into(),
828                    component: "comp".into(),
829                    endpoint: "generate".into(),
830                    card_json: serde_json::json!({
831                        "display_name": "base",
832                        "aliases": ["alias"],
833                        "source_path": "org/base",
834                    }),
835                    model_suffix: None,
836                };
837                let adapter = lora_model_spec(
838                    "ns",
839                    "comp",
840                    "generate",
841                    adapter_name,
842                    "org/base",
843                    adapter_name,
844                );
845                let (incumbent, newcomer) = if adapter_first {
846                    (adapter, base)
847                } else {
848                    (base, adapter)
849                };
850                let incumbent = first.register(incumbent).await.unwrap();
851                let result = second.register(newcomer).await;
852                assert_eq!(
853                    result.is_ok(),
854                    compatible,
855                    "{adapter_name}, adapter_first={adapter_first}: {result:?}"
856                );
857                let instances = first
858                    .list(DiscoveryQuery::EndpointModels {
859                        namespace: "ns".into(),
860                        component: "comp".into(),
861                        endpoint: "generate".into(),
862                    })
863                    .await
864                    .unwrap();
865                assert!(instances.contains(&incumbent));
866                assert_eq!(instances.len(), if compatible { 2 } else { 1 });
867            }
868        }
869    }
870
871    #[tokio::test]
872    async fn register_rejects_lora_adapter_for_different_base_model() {
873        let registry = SharedMockRegistry::new();
874        let discovery1 = MockDiscovery::new(Some(1), registry.clone());
875        let discovery2 = MockDiscovery::new(Some(2), registry);
876
877        discovery1
878            .register(DiscoverySpec::Model {
879                namespace: "ns".to_string(),
880                component: "comp".to_string(),
881                endpoint: "generate".to_string(),
882                card_json: serde_json::json!({
883                    "display_name": "base-model",
884                    "source_path": "base-repo",
885                }),
886                model_suffix: None,
887            })
888            .await
889            .unwrap();
890
891        let err = discovery2
892            .register(lora_model_spec(
893                "ns",
894                "comp",
895                "generate",
896                "adapter-a",
897                "other-base-repo",
898                "adapter-a",
899            ))
900            .await
901            .unwrap_err();
902
903        assert!(err.to_string().contains(
904            "Cannot register model 'adapter-a' on endpoint 'ns/comp/generate': a different model 'base-model' is already registered there"
905        ));
906    }
907}