1use 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#[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
28pub 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
50fn matches_query(instance: &DiscoveryInstance, query: &DiscoveryQuery) -> bool {
52 match (instance, query) {
53 (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 (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 (
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 (
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 let mut stream = client1.list_and_watch(query.clone(), None).await.unwrap();
519
520 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 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 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 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}