1use 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#[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
25pub 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
47fn matches_query(instance: &DiscoveryInstance, query: &DiscoveryQuery) -> bool {
49 match (instance, query) {
50 (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 (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 (
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 (
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 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 for id in known_instances.difference(¤t_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 let mut stream = client1.list_and_watch(query.clone(), None).await.unwrap();
336
337 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 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 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}