use anyhow::Result;
use k8s_openapi::apimachinery::pkg::apis::meta::v1::OwnerReference;
use kube::{
Api, Client as KubeClient, CustomResource,
api::{Patch, PatchParams},
};
use serde::{Deserialize, Serialize};
use crate::discovery::{DiscoveryMetadata, EventScope};
const FIELD_MANAGER: &str = "dynamo-worker";
#[derive(CustomResource, Clone, Debug, Deserialize, Serialize)]
#[kube(
group = "nvidia.com",
version = "v1alpha1",
kind = "DynamoWorkerMetadata",
namespaced,
schema = "disabled"
)]
pub struct DynamoWorkerMetadataSpec {
pub data: serde_json::Value,
}
impl DynamoWorkerMetadataSpec {
pub fn new(data: serde_json::Value) -> Self {
Self { data }
}
}
pub fn build_cr(
cr_name: &str,
pod_name: &str,
pod_uid: &str,
metadata: &DiscoveryMetadata,
) -> Result<DynamoWorkerMetadata> {
let mut data = serde_json::to_value(metadata)?;
add_legacy_event_channel_fields(&mut data)?;
let spec = DynamoWorkerMetadataSpec::new(data);
let mut cr = DynamoWorkerMetadata::new(cr_name, spec);
cr.metadata.owner_references = Some(vec![OwnerReference {
api_version: "v1".to_string(),
kind: "Pod".to_string(),
name: pod_name.to_string(),
uid: pod_uid.to_string(),
controller: Some(cr_name == pod_name),
block_owner_deletion: Some(false),
}]);
Ok(cr)
}
pub(super) fn deserialize_metadata(mut data: serde_json::Value) -> Result<DiscoveryMetadata> {
add_current_event_channel_scopes(&mut data)?;
Ok(serde_json::from_value(data)?)
}
fn add_legacy_event_channel_fields(data: &mut serde_json::Value) -> Result<()> {
let Some(channels) = data
.get_mut("event_channels")
.and_then(serde_json::Value::as_object_mut)
else {
return Ok(());
};
for channel in channels.values_mut() {
let Some(channel) = channel.as_object_mut() else {
continue;
};
let Some(scope) = channel.get("scope").cloned() else {
continue;
};
let scope = serde_json::from_value::<EventScope>(scope)?;
channel.insert(
"namespace".to_string(),
serde_json::Value::String(scope.namespace().to_string()),
);
channel.insert(
"component".to_string(),
serde_json::Value::String(scope.component().unwrap_or("").to_string()),
);
}
Ok(())
}
fn add_current_event_channel_scopes(data: &mut serde_json::Value) -> Result<()> {
let Some(channels) = data
.get_mut("event_channels")
.and_then(serde_json::Value::as_object_mut)
else {
return Ok(());
};
for channel in channels.values_mut() {
let Some(channel) = channel.as_object_mut() else {
continue;
};
if channel.contains_key("scope") {
continue;
}
let (Some(namespace), Some(component)) = (
channel.get("namespace").and_then(serde_json::Value::as_str),
channel.get("component").and_then(serde_json::Value::as_str),
) else {
continue;
};
let scope = if component.is_empty() {
EventScope::Namespace {
name: namespace.to_string(),
}
} else {
EventScope::Component {
namespace: namespace.to_string(),
component: component.to_string(),
}
};
channel.insert("scope".to_string(), serde_json::to_value(scope)?);
}
Ok(())
}
pub async fn apply_cr(
kube_client: &KubeClient,
namespace: &str,
cr: &DynamoWorkerMetadata,
) -> Result<()> {
let api: Api<DynamoWorkerMetadata> = Api::namespaced(kube_client.clone(), namespace);
let cr_name = cr
.metadata
.name
.as_ref()
.ok_or_else(|| anyhow::anyhow!("CR must have a name"))?;
let params = PatchParams::apply(FIELD_MANAGER).force();
api.patch(cr_name, ¶ms, &Patch::Apply(cr))
.await
.map_err(|e| anyhow::anyhow!("Failed to apply DynamoWorkerMetadata CR: {}", e))?;
tracing::debug!(
"Applied DynamoWorkerMetadata CR: name={}, namespace={}",
cr_name,
namespace
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::discovery::{
DiscoveryInstance, DiscoveryQuery, EventChannelQuery, EventScope, EventSourceQuery,
EventTransport, MAX_JSON_SAFE_PUBLISHER_ID,
};
use crate::protocols::EndpointId;
use kube::Resource;
#[test]
fn test_crd_metadata() {
assert_eq!(DynamoWorkerMetadata::group(&()), "nvidia.com");
assert_eq!(DynamoWorkerMetadata::version(&()), "v1alpha1");
assert_eq!(DynamoWorkerMetadata::kind(&()), "DynamoWorkerMetadata");
assert_eq!(DynamoWorkerMetadata::plural(&()), "dynamoworkermetadatas");
}
#[test]
fn test_serialization_roundtrip() {
let data = serde_json::json!({
"endpoints": {
"ns/comp/ep": {
"type": "Endpoint",
"namespace": "ns",
"component": "comp",
"endpoint": "ep",
"instance_id": 12345,
"transport": { "Nats": "nats://localhost:4222" }
}
},
"model_cards": {}
});
let spec = DynamoWorkerMetadataSpec::new(data.clone());
let cr = DynamoWorkerMetadata::new("test-pod", spec);
let json = serde_json::to_string(&cr).expect("Failed to serialize CR");
let deserialized: DynamoWorkerMetadata =
serde_json::from_str(&json).expect("Failed to deserialize CR");
assert_eq!(deserialized.spec.data, data);
}
#[test]
fn event_source_round_trips_through_kubernetes_metadata() {
let mut metadata = DiscoveryMetadata::new();
let endpoint = EndpointId {
namespace: "workers".to_string(),
component: "backend".to_string(),
name: "kv-state".to_string(),
};
let source = DiscoveryInstance::EventSource {
scope: EventScope::Endpoint {
endpoint: endpoint.clone(),
},
topic: "kv-events".to_string(),
publisher_id: MAX_JSON_SAFE_PUBLISHER_ID,
metadata: serde_json::json!({"worker_id": 7, "dp_rank": 0}),
};
metadata.register_event_source(source.clone()).unwrap();
let cr = build_cr("test-pod", "test-pod", "pod-uid", &metadata).unwrap();
let publisher_id = cr.spec.data["event_sources"]
.as_object()
.and_then(|sources| sources.values().next())
.and_then(|source| source.get("publisher_id"))
.expect("serialized event source publisher ID");
assert_eq!(publisher_id.as_u64(), Some(MAX_JSON_SAFE_PUBLISHER_ID));
let round_trip = deserialize_metadata(cr.spec.data).unwrap();
assert_eq!(
round_trip.filter(&DiscoveryQuery::EventSources(
EventSourceQuery::endpoint_topic(endpoint, "kv-events")
)),
vec![source]
);
}
#[test]
fn event_channel_metadata_supports_v1_2_wire_shape() {
#[derive(serde::Deserialize)]
#[serde(tag = "type")]
enum LegacyDiscoveryInstance {
EventChannel {
namespace: String,
component: String,
topic: String,
instance_id: u64,
transport: EventTransport,
},
}
let endpoint = EndpointId {
namespace: "workers".to_string(),
component: "backend".to_string(),
name: "generate".to_string(),
};
let transport = EventTransport::zmq("tcp://worker:5555");
let channel = DiscoveryInstance::EventChannel {
scope: EventScope::Endpoint {
endpoint: endpoint.clone(),
},
topic: "kv-events".to_string(),
instance_id: 42,
transport: transport.clone(),
};
let mut metadata = DiscoveryMetadata::new();
metadata.register_event_channel(channel.clone()).unwrap();
let cr = build_cr("test-pod", "test-pod", "pod-uid", &metadata).unwrap();
let encoded_channel = cr.spec.data["event_channels"]
.as_object()
.and_then(|channels| channels.values().next())
.cloned()
.expect("serialized event channel");
let LegacyDiscoveryInstance::EventChannel {
namespace,
component,
topic,
instance_id,
transport: legacy_transport,
} = serde_json::from_value(encoded_channel).expect("v1.2 event channel shape");
assert_eq!(
(namespace, component, topic, instance_id, legacy_transport),
(
"workers".to_string(),
"backend".to_string(),
"kv-events".to_string(),
42,
transport.clone(),
)
);
let current_round_trip = deserialize_metadata(cr.spec.data).unwrap();
assert_eq!(
current_round_trip.filter(&DiscoveryQuery::EventChannels(
EventChannelQuery::endpoint_topic(endpoint, "kv-events")
)),
vec![channel]
);
let namespace_transport = EventTransport::nats("namespace.workers");
let legacy_metadata = serde_json::json!({
"endpoints": {},
"model_cards": {},
"event_channels": {
"workers/backend/kv-events/2a": {
"type": "EventChannel",
"namespace": "workers",
"component": "backend",
"topic": "kv-events",
"instance_id": 42,
"transport": transport,
},
"workers//namespace-events/2b": {
"type": "EventChannel",
"namespace": "workers",
"component": "",
"topic": "namespace-events",
"instance_id": 43,
"transport": namespace_transport,
}
}
});
let upgraded = deserialize_metadata(legacy_metadata).unwrap();
assert_eq!(
upgraded.filter(&DiscoveryQuery::EventChannels(EventChannelQuery::topic(
"workers",
"backend",
"kv-events",
))),
vec![DiscoveryInstance::EventChannel {
scope: EventScope::Component {
namespace: "workers".to_string(),
component: "backend".to_string(),
},
topic: "kv-events".to_string(),
instance_id: 42,
transport,
}]
);
assert_eq!(
upgraded.filter(&DiscoveryQuery::EventChannels(
EventChannelQuery::namespace_topic("workers", "namespace-events")
)),
vec![DiscoveryInstance::EventChannel {
scope: EventScope::Namespace {
name: "workers".to_string(),
},
topic: "namespace-events".to_string(),
instance_id: 43,
transport: namespace_transport,
}]
);
}
}