1mod crd;
5mod daemon;
6mod utils;
7
8pub use crd::{DynamoWorkerMetadata, DynamoWorkerMetadataSpec};
9pub use utils::{hash_container_name, hash_pod_name};
12
13use crd::{apply_cr, build_cr};
14use daemon::DiscoveryDaemon;
15use utils::{KubeDiscoveryMode, PodInfo};
16
17use crate::CancellationToken;
18use crate::discovery::{
19 Discovery, DiscoveryEvent, DiscoveryInstance, DiscoveryInstanceId, DiscoveryMetadata,
20 DiscoveryQuery, DiscoverySpec, DiscoveryStream, MAX_JSON_SAFE_PUBLISHER_ID, MetadataSnapshot,
21 ModelCardInstanceId, reconcile_discovery_snapshot,
22};
23use anyhow::Result;
24use async_trait::async_trait;
25use kube::{Api, Client as KubeClient, api::DeleteParams};
26use std::collections::{HashMap, HashSet};
27use std::future::Future;
28use std::sync::Arc;
29use tokio::sync::RwLock;
30
31fn validate_kubernetes_publisher_id(publisher_id: u64) -> Result<()> {
32 if publisher_id > MAX_JSON_SAFE_PUBLISHER_ID {
33 anyhow::bail!(
34 "Kubernetes discovery publisher ID {publisher_id} exceeds the JSON-safe maximum \
35 {MAX_JSON_SAFE_PUBLISHER_ID}"
36 );
37 }
38
39 Ok(())
40}
41
42async fn update_model_taints_and_persist<F, Fut>(
43 metadata: &Arc<RwLock<DiscoveryMetadata>>,
44 id: ModelCardInstanceId,
45 taints: HashSet<String>,
46 persist: F,
47) -> Result<bool>
48where
49 F: FnOnce(DiscoveryMetadata) -> Fut + Send + 'static,
50 Fut: Future<Output = Result<DiscoveryMetadata>> + Send + 'static,
51{
52 let metadata = Arc::clone(metadata);
53 tokio::spawn(async move {
57 let mut metadata = metadata.write().await;
58 let mut candidate = metadata.clone();
59 let changed = candidate.update_model_taints(&id, taints)?;
60
61 let persisted = persist(candidate).await?;
64 *metadata = persisted;
65 Ok(changed)
66 })
67 .await
68 .map_err(|error| anyhow::anyhow!("model taint persistence task failed: {error}"))?
69}
70
71#[derive(Clone)]
73pub struct KubeDiscoveryClient {
74 instance_id: u64,
75 metadata: Arc<RwLock<DiscoveryMetadata>>,
76 metadata_watch: tokio::sync::watch::Receiver<Arc<MetadataSnapshot>>,
77 kube_client: KubeClient,
78 pod_info: PodInfo,
79}
80
81impl KubeDiscoveryClient {
82 pub async fn new(
88 metadata: Arc<RwLock<DiscoveryMetadata>>,
89 cancel_token: CancellationToken,
90 ) -> Result<Self> {
91 let pod_info = PodInfo::from_env()?;
92 let instance_id = pod_info.target.instance_id();
93 let cr_name = pod_info.target.cr_name();
94
95 tracing::info!(
96 "Initializing KubeDiscoveryClient: mode={:?}, target={:?}, cr_name={}, instance_id={:x}, namespace={}, pod_uid={}",
97 pod_info.mode,
98 pod_info.target,
99 cr_name,
100 instance_id,
101 pod_info.pod_namespace,
102 pod_info.pod_uid
103 );
104
105 let kube_client = KubeClient::try_default()
106 .await
107 .map_err(|e| anyhow::anyhow!("Failed to create Kubernetes client: {}", e))?;
108
109 if pod_info.mode == KubeDiscoveryMode::Container {
114 let cr_api: Api<DynamoWorkerMetadata> =
115 Api::namespaced(kube_client.clone(), &pod_info.pod_namespace);
116 match cr_api.delete(&cr_name, &DeleteParams::default()).await {
117 Ok(_) => tracing::info!("Deleted stale CR: {}", cr_name),
118 Err(kube::Error::Api(err_resp)) if err_resp.code == 404 => {
119 tracing::debug!("No stale CR to delete: {}", cr_name);
120 }
121 Err(e) => {
122 panic!(
123 "Failed to clear stale CR '{}': {} — cannot start with stale discovery state",
124 cr_name, e
125 );
126 }
127 }
128 }
129
130 let (watch_tx, watch_rx) = tokio::sync::watch::channel(Arc::new(MetadataSnapshot::empty()));
132
133 let daemon = DiscoveryDaemon::new(kube_client.clone(), pod_info.clone(), cancel_token)?;
135
136 tokio::spawn(async move {
137 if let Err(e) = daemon.run(watch_tx).await {
138 tracing::error!("Discovery daemon failed: {e}");
139 }
140 });
141
142 tracing::info!("Discovery daemon started");
143
144 Ok(Self {
145 instance_id,
146 metadata,
147 metadata_watch: watch_rx,
148 kube_client,
149 pod_info,
150 })
151 }
152}
153
154#[async_trait]
155impl Discovery for KubeDiscoveryClient {
156 fn instance_id(&self) -> u64 {
157 self.instance_id
158 }
159
160 async fn register_internal(&self, spec: DiscoverySpec) -> Result<DiscoveryInstance> {
161 match &spec {
162 DiscoverySpec::EventChannel { publisher_id, .. }
163 | DiscoverySpec::EventSource { publisher_id, .. } => {
164 validate_kubernetes_publisher_id(*publisher_id)?;
165 }
166 _ => {}
167 }
168 let instance = spec.into_instance(self.instance_id());
169 let instance_id = instance.instance_id();
170
171 tracing::debug!(
172 "Registering discovery instance: {:?}, instance_id={:x}",
173 instance,
174 instance_id
175 );
176
177 let mut metadata = self.metadata.write().await;
180
181 let original_state = metadata.clone();
183
184 let registered_instance = match &instance {
185 DiscoveryInstance::Endpoint(inst) => {
186 tracing::info!(
187 "Registering endpoint: namespace={}, component={}, endpoint={}, instance_id={:x}",
188 inst.namespace,
189 inst.component,
190 inst.endpoint,
191 instance_id
192 );
193 metadata.register_endpoint(instance.clone())?;
194 instance.clone()
195 }
196 DiscoveryInstance::Model {
197 namespace,
198 component,
199 endpoint,
200 ..
201 } => {
202 tracing::info!(
203 "Registering model card: namespace={}, component={}, endpoint={}, instance_id={:x}",
204 namespace,
205 component,
206 endpoint,
207 instance_id
208 );
209 metadata.register_model_card(instance.clone())?
210 }
211 DiscoveryInstance::EventChannel { scope, topic, .. } => {
212 tracing::info!(
213 "Registering event channel: scope={:?}, topic={}, instance_id={:x}",
214 scope,
215 topic,
216 instance_id
217 );
218 metadata.register_event_channel(instance.clone())?;
219 instance.clone()
220 }
221 DiscoveryInstance::EventSource { scope, topic, .. } => {
222 tracing::info!(
223 "Registering event source: scope={:?}, topic={}, publisher_id={:x}",
224 scope,
225 topic,
226 instance_id
227 );
228 metadata.register_event_source(instance.clone())?;
229 instance.clone()
230 }
231 };
232
233 let cr_name = self.pod_info.target.cr_name();
236 let cr = build_cr(
237 &cr_name,
238 &self.pod_info.pod_name,
239 &self.pod_info.pod_uid,
240 &metadata,
241 )?;
242
243 if let Err(e) = apply_cr(&self.kube_client, &self.pod_info.pod_namespace, &cr).await {
244 tracing::warn!(
246 "Failed to persist metadata to CR, rolling back local state: {}",
247 e
248 );
249 *metadata = original_state;
250 return Err(e);
251 }
252
253 tracing::debug!("Persisted metadata to DynamoWorkerMetadata CR");
254
255 Ok(registered_instance)
256 }
257
258 async fn update_model_taints_internal(
259 &self,
260 id: ModelCardInstanceId,
261 taints: HashSet<String>,
262 ) -> Result<()> {
263 let kube_client = self.kube_client.clone();
264 let pod_namespace = self.pod_info.pod_namespace.clone();
265 let cr_name = self.pod_info.target.cr_name();
266 let pod_name = self.pod_info.pod_name.clone();
267 let pod_uid = self.pod_info.pod_uid.clone();
268 let changed = update_model_taints_and_persist(
269 &self.metadata,
270 id,
271 taints,
272 move |candidate| async move {
273 let cr = build_cr(&cr_name, &pod_name, &pod_uid, &candidate)?;
274 apply_cr(&kube_client, &pod_namespace, &cr).await?;
275 Ok(candidate)
276 },
277 )
278 .await?;
279 if !changed {
280 return Ok(());
281 }
282
283 tracing::debug!("Persisted model taint update to DynamoWorkerMetadata CR");
284 Ok(())
285 }
286
287 async fn unregister(&self, instance: DiscoveryInstance) -> Result<()> {
288 let instance_id = instance.instance_id();
289
290 let mut metadata = self.metadata.write().await;
293
294 let original_state = metadata.clone();
296
297 match &instance {
298 DiscoveryInstance::Endpoint(inst) => {
299 tracing::info!(
300 "Unregistering endpoint: namespace={}, component={}, endpoint={}, instance_id={:x}",
301 inst.namespace,
302 inst.component,
303 inst.endpoint,
304 instance_id
305 );
306 metadata.unregister_endpoint(&instance)?;
307 }
308 DiscoveryInstance::Model {
309 namespace,
310 component,
311 endpoint,
312 ..
313 } => {
314 tracing::info!(
315 "Unregistering model card: namespace={}, component={}, endpoint={}, instance_id={:x}",
316 namespace,
317 component,
318 endpoint,
319 instance_id
320 );
321 metadata.unregister_model_card(&instance)?;
322 }
323 DiscoveryInstance::EventChannel { scope, topic, .. } => {
324 tracing::info!(
325 "Unregistering event channel: scope={:?}, topic={}, instance_id={:x}",
326 scope,
327 topic,
328 instance_id
329 );
330 metadata.unregister_event_channel(&instance)?;
331 }
332 DiscoveryInstance::EventSource { scope, topic, .. } => {
333 tracing::info!(
334 "Unregistering event source: scope={:?}, topic={}, publisher_id={:x}",
335 scope,
336 topic,
337 instance_id
338 );
339 metadata.unregister_event_source(&instance)?;
340 }
341 }
342
343 let cr_name = self.pod_info.target.cr_name();
346 let cr = build_cr(
347 &cr_name,
348 &self.pod_info.pod_name,
349 &self.pod_info.pod_uid,
350 &metadata,
351 )?;
352
353 if let Err(e) = apply_cr(&self.kube_client, &self.pod_info.pod_namespace, &cr).await {
354 tracing::warn!(
356 "Failed to persist metadata removal to CR, rolling back local state: {}",
357 e
358 );
359 *metadata = original_state;
360 return Err(e);
361 }
362
363 tracing::debug!("Persisted metadata removal to DynamoWorkerMetadata CR");
364
365 Ok(())
366 }
367
368 async fn list(&self, query: DiscoveryQuery) -> Result<Vec<DiscoveryInstance>> {
369 tracing::debug!("KubeDiscoveryClient::list called with query={:?}", query);
370
371 let snapshot = self.metadata_watch.borrow().clone();
373
374 tracing::debug!(
375 "List using snapshot seq={} with {} instances",
376 snapshot.sequence,
377 snapshot.instances.len()
378 );
379
380 let instances = snapshot.filter(&query);
382
383 tracing::info!(
384 "KubeDiscoveryClient::list returning {} instances for query={:?}",
385 instances.len(),
386 query
387 );
388
389 Ok(instances)
390 }
391
392 async fn list_and_watch(
393 &self,
394 query: DiscoveryQuery,
395 cancel_token: Option<CancellationToken>,
396 ) -> Result<DiscoveryStream> {
397 use tokio::sync::mpsc;
398
399 tracing::info!(
400 "KubeDiscoveryClient::list_and_watch started for query={:?}",
401 query
402 );
403
404 let mut watch_rx = self.metadata_watch.clone();
406
407 let (event_tx, event_rx) = mpsc::unbounded_channel();
409
410 let stream_id = uuid::Uuid::new_v4();
412
413 tokio::spawn(async move {
415 let initial_snapshot = watch_rx.borrow_and_update().clone();
419
420 let initial: HashMap<DiscoveryInstanceId, DiscoveryInstance> = initial_snapshot
422 .instances
423 .values()
424 .flat_map(|metadata| metadata.filter(&query))
425 .map(|instance| (instance.id(), instance))
426 .collect();
427
428 tracing::debug!(
429 stream_id = %stream_id,
430 initial_count = initial.len(),
431 "Watch started for query={:?}",
432 query
433 );
434
435 for instance in initial.values() {
437 tracing::info!(
438 stream_id = %stream_id,
439 instance_id = format!("{:x}", instance.instance_id()),
440 "Emitting initial Added event"
441 );
442 if event_tx
443 .send(Ok(DiscoveryEvent::Added(instance.clone())))
444 .is_err()
445 {
446 tracing::debug!(
447 stream_id = %stream_id,
448 "Watch receiver dropped during initial sync"
449 );
450 return;
451 }
452 }
453
454 let mut known = initial;
456
457 loop {
458 tracing::trace!(
459 stream_id = %stream_id,
460 known_count = known.len(),
461 "Watch loop waiting for changes"
462 );
463
464 let watch_result = if let Some(ref token) = cancel_token {
466 tokio::select! {
467 result = watch_rx.changed() => result,
468 _ = token.cancelled() => {
469 tracing::info!(
470 stream_id = %stream_id,
471 "Watch cancelled via cancel token"
472 );
473 break;
474 }
475 }
476 } else {
477 watch_rx.changed().await
478 };
479
480 match watch_result {
481 Ok(()) => {
482 let snapshot = watch_rx.borrow_and_update().clone();
484
485 let current: HashMap<DiscoveryInstanceId, DiscoveryInstance> = snapshot
487 .instances
488 .values()
489 .flat_map(|metadata| metadata.filter(&query))
490 .map(|instance| (instance.id(), instance))
491 .collect();
492
493 tracing::debug!(
494 stream_id = %stream_id,
495 seq = snapshot.sequence,
496 current_count = current.len(),
497 known_count = known.len(),
498 "Watch received snapshot update"
499 );
500
501 let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
502
503 if events.is_empty() {
505 tracing::debug!(
506 stream_id = %stream_id,
507 seq = snapshot.sequence,
508 "Watch snapshot received but no diff detected"
509 );
510 } else {
511 tracing::debug!(
512 stream_id = %stream_id,
513 seq = snapshot.sequence,
514 emitted_events = events.len(),
515 total = reconciled.len(),
516 "Watch detected changes"
517 );
518 }
519
520 for event in events {
521 let (event_kind, instance_id) = match &event {
522 DiscoveryEvent::Added(instance) => ("added", instance.id()),
523 DiscoveryEvent::ModelTaintsUpdated(update) => (
524 "model_taints_updated",
525 DiscoveryInstanceId::Model(update.id.clone()),
526 ),
527 DiscoveryEvent::Removed(id) => ("removed", id.clone()),
528 };
529 tracing::info!(
530 stream_id = %stream_id,
531 event_kind,
532 ?instance_id,
533 "Emitting discovery event"
534 );
535 tracing::debug!(
536 stream_id = %stream_id,
537 ?event,
538 "Discovery event detail"
539 );
540 if event_tx.send(Ok(event)).is_err() {
541 tracing::debug!(stream_id = %stream_id, "Watch receiver dropped");
542 return;
543 }
544 }
545
546 known = reconciled;
547 }
548 Err(_) => {
549 tracing::info!(
550 stream_id = %stream_id,
551 "Watch channel closed (daemon stopped)"
552 );
553 break;
554 }
555 }
556 }
557 });
558
559 let stream = tokio_stream::wrappers::UnboundedReceiverStream::new(event_rx);
561 Ok(Box::pin(stream))
562 }
563}
564
565#[cfg(test)]
566mod tests {
567 use super::*;
568 use crate::component::TransportType;
569 use crate::discovery::{EventScope, EventTransport, ModelTaintsUpdate};
570
571 fn endpoint_instance(instance_id: u64, transport: &str) -> DiscoveryInstance {
572 DiscoveryInstance::Endpoint(crate::component::Instance {
573 namespace: "ns".to_string(),
574 component: "component".to_string(),
575 endpoint: "endpoint".to_string(),
576 instance_id,
577 transport: TransportType::Tcp(transport.to_string()),
578 device_type: None,
579 request_plane_codec: None,
580 })
581 }
582
583 fn model_with_taint(taint: &str) -> DiscoveryInstance {
584 DiscoveryInstance::Model {
585 namespace: "ns".to_string(),
586 component: "worker".to_string(),
587 endpoint: "generate".to_string(),
588 instance_id: 7,
589 card_json: serde_json::json!({
590 "runtime_config": {"taints": [taint]}
591 }),
592 model_suffix: None,
593 }
594 }
595
596 #[test]
597 fn publisher_ids_must_fit_kubernetes_json_safe_range() {
598 assert!(validate_kubernetes_publisher_id(MAX_JSON_SAFE_PUBLISHER_ID).is_ok());
599 assert!(validate_kubernetes_publisher_id(MAX_JSON_SAFE_PUBLISHER_ID + 1).is_err());
600 assert!(validate_kubernetes_publisher_id(u64::MAX).is_err());
601 }
602
603 #[test]
604 fn snapshot_diff_emits_updated_instance_when_transport_changes() {
605 let original = endpoint_instance(1, "127.0.0.1:8000");
606 let updated = endpoint_instance(1, "127.0.0.1:9000");
607 let known = HashMap::from([(original.id(), original)]);
608 let current = HashMap::from([(updated.id(), updated.clone())]);
609
610 let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
611
612 assert_eq!(events, vec![DiscoveryEvent::Added(updated.clone())]);
613 assert_eq!(reconciled.get(&updated.id()), Some(&updated));
614 }
615
616 #[test]
617 fn snapshot_diff_ignores_same_id_event_channel_changes() {
618 let event_channel = |endpoint: &str| DiscoveryInstance::EventChannel {
619 scope: EventScope::Namespace {
620 name: "ns".to_string(),
621 },
622 topic: "topic".to_string(),
623 instance_id: 1,
624 transport: EventTransport::zmq(endpoint),
625 };
626 let original = event_channel("tcp://127.0.0.1:8000");
627 let updated = event_channel("tcp://127.0.0.1:9000");
628 let known = HashMap::from([(original.id(), original.clone())]);
629 let current = HashMap::from([(updated.id(), updated)]);
630
631 let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
632
633 assert!(events.is_empty());
634 assert_eq!(reconciled.get(&original.id()), Some(&original));
635 }
636
637 #[test]
638 fn snapshot_diff_emits_added_and_removed_instances() {
639 let removed_instance = endpoint_instance(1, "127.0.0.1:8000");
640 let added_instance = endpoint_instance(2, "127.0.0.1:9000");
641 let removed_id = removed_instance.id();
642 let added_id = added_instance.id();
643 let known = HashMap::from([(removed_id.clone(), removed_instance)]);
644 let current = HashMap::from([(added_id.clone(), added_instance.clone())]);
645
646 let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
647
648 assert_eq!(events.len(), 2);
649 assert!(events.contains(&DiscoveryEvent::Removed(removed_id.clone())));
650 assert!(events.contains(&DiscoveryEvent::Added(added_instance.clone())));
651 assert!(!reconciled.contains_key(&removed_id));
652 assert_eq!(reconciled.get(&added_id), Some(&added_instance));
653 }
654 #[test]
655 fn changed_model_taints_emit_scoped_event() {
656 let old = model_with_taint("old");
657 let updated = model_with_taint("updated");
658 let known = HashMap::from([(old.id(), old)]);
659 let current = HashMap::from([(updated.id(), updated.clone())]);
660
661 let (events, reconciled) = reconcile_discovery_snapshot(&known, current);
662
663 let DiscoveryInstanceId::Model(id) = updated.id() else {
664 unreachable!()
665 };
666 assert_eq!(
667 events,
668 vec![DiscoveryEvent::ModelTaintsUpdated(ModelTaintsUpdate {
669 id,
670 taints: vec!["updated".to_string()],
671 })]
672 );
673 assert_eq!(reconciled.get(&updated.id()), Some(&updated));
674 }
675
676 #[tokio::test]
677 async fn model_taint_persistence_completes_after_caller_cancellation() {
678 let model = model_with_taint("old");
679 let DiscoveryInstanceId::Model(id) = model.id() else {
680 unreachable!()
681 };
682 let mut initial = DiscoveryMetadata::new();
683 initial.register_model_card(model).unwrap();
684 let metadata = Arc::new(RwLock::new(initial));
685 let task_metadata = metadata.clone();
686 let remote = Arc::new(RwLock::new(DiscoveryMetadata::new()));
687 let task_remote = remote.clone();
688 let (remote_committed_tx, remote_committed_rx) = tokio::sync::oneshot::channel();
689 let (ack_tx, ack_rx) = tokio::sync::oneshot::channel();
690
691 let task = tokio::spawn(async move {
692 update_model_taints_and_persist(
693 &task_metadata,
694 id,
695 HashSet::from(["new".to_string()]),
696 move |candidate| async move {
697 *task_remote.write().await = candidate.clone();
698 remote_committed_tx.send(()).unwrap();
699 ack_rx.await.unwrap();
700 Ok(candidate)
701 },
702 )
703 .await
704 });
705
706 remote_committed_rx.await.unwrap();
707 task.abort();
708 assert!(task.await.unwrap_err().is_cancelled());
709 ack_tx.send(()).unwrap();
710
711 let stored = tokio::time::timeout(std::time::Duration::from_secs(1), async {
712 loop {
713 let stored = metadata.read().await.get_all_model_cards().pop().unwrap();
714 let DiscoveryInstance::Model { card_json, .. } = &stored else {
715 unreachable!()
716 };
717 if card_json["runtime_config"]["taints"] == serde_json::json!(["new"]) {
718 break stored;
719 }
720 tokio::task::yield_now().await;
721 }
722 })
723 .await
724 .expect("detached persistence did not commit local metadata");
725 let DiscoveryInstance::Model { card_json, .. } = stored else {
726 unreachable!()
727 };
728 assert_eq!(
729 card_json["runtime_config"]["taints"],
730 serde_json::json!(["new"])
731 );
732 let remote = remote.read().await.get_all_model_cards().pop().unwrap();
733 let DiscoveryInstance::Model { card_json, .. } = remote else {
734 unreachable!()
735 };
736 assert_eq!(
737 card_json["runtime_config"]["taints"],
738 serde_json::json!(["new"])
739 );
740 }
741
742 #[tokio::test]
743 async fn local_noop_reapplies_authoritative_model_taints() {
744 let local_model = model_with_taint("old");
745 let DiscoveryInstanceId::Model(id) = local_model.id() else {
746 unreachable!()
747 };
748 let mut initial = DiscoveryMetadata::new();
749 initial.register_model_card(local_model).unwrap();
750 let metadata = Arc::new(RwLock::new(initial));
751 let persisted = Arc::new(RwLock::new(None));
752 let task_persisted = persisted.clone();
753
754 let changed = update_model_taints_and_persist(
755 &metadata,
756 id,
757 HashSet::from(["old".to_string()]),
758 move |candidate| async move {
759 *task_persisted.write().await = Some(candidate.clone());
760 Ok(candidate)
761 },
762 )
763 .await
764 .unwrap();
765
766 assert!(!changed);
767 let reapplied = persisted
768 .read()
769 .await
770 .clone()
771 .expect("no-op was not persisted");
772 let DiscoveryInstance::Model { card_json, .. } =
773 reapplied.get_all_model_cards().pop().unwrap()
774 else {
775 unreachable!()
776 };
777 assert_eq!(
778 card_json["runtime_config"]["taints"],
779 serde_json::json!(["old"])
780 );
781 }
782}