use anyhow::Result;
use async_trait::async_trait;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use uuid::Uuid;
use crate::runtime::{Endpoint, GearInstance, GearManager, GrpcServiceNameConflict};
fn openapi_spec_hash(spec: &str) -> String {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
spec.hash(&mut hasher);
format!("{:016x}", hasher.finish())
}
pub use cf_system_sdks::directory::{
DirectoryClient, DirectoryInvalidArgument, DirectoryNotFound, DirectoryPermissionDenied,
DirectoryServiceNameConflict, GrpcServiceInfo, InstanceState, LabelSelector,
RegisterInstanceInfo, ServiceEndpoint, ServiceInstanceInfo,
};
fn runtime_state_to_domain(state: crate::runtime::InstanceState) -> InstanceState {
use crate::runtime::InstanceState as Rt;
match state {
Rt::Registered => InstanceState::Registered,
Rt::Ready => InstanceState::Ready,
Rt::Healthy => InstanceState::Healthy,
Rt::Quarantined => InstanceState::Quarantined,
Rt::Draining => InstanceState::Draining,
}
}
fn project_instance(gear: &str, inst: &GearInstance) -> ServiceInstanceInfo {
let endpoint = inst
.grpc_services
.values()
.next()
.or(inst.rest_endpoint.as_ref())
.map(|ep| ServiceEndpoint::new(ep.uri.clone()));
ServiceInstanceInfo::new(gear, inst.instance_id.to_string())
.with_endpoint(endpoint)
.with_version(inst.version.clone())
.with_rest_endpoint(
inst.rest_endpoint
.as_ref()
.map(|ep| ServiceEndpoint::new(ep.uri.clone())),
)
.with_openapi_spec_hash(inst.openapi_spec.as_deref().map(openapi_spec_hash))
.with_grpc_services(
inst.grpc_services
.iter()
.map(|(name, e)| (name.clone(), ServiceEndpoint::new(e.uri.clone())))
.collect(),
)
.with_labels(inst.labels.clone())
.with_state(runtime_state_to_domain(inst.state()))
}
pub struct LocalDirectoryClient {
mgr: Arc<GearManager>,
}
impl LocalDirectoryClient {
#[must_use]
pub fn new(mgr: Arc<GearManager>) -> Self {
Self { mgr }
}
}
#[async_trait]
impl DirectoryClient for LocalDirectoryClient {
async fn resolve_grpc_service(&self, service_name: &str) -> Result<ServiceEndpoint> {
if let Some((_gear, _inst, ep)) = self.mgr.pick_service_round_robin(service_name) {
return Ok(ServiceEndpoint::new(ep.uri));
}
Err(DirectoryNotFound::new(format!("service {service_name}")).into())
}
async fn resolve_rest_service(&self, gear_name: &str) -> Result<ServiceEndpoint> {
if let Some(ep) = self.mgr.pick_rest_endpoint_round_robin(gear_name) {
return Ok(ServiceEndpoint::new(ep.uri));
}
Err(DirectoryNotFound::new(format!("gear {gear_name}")).into())
}
async fn get_openapi_spec(&self, gear_name: &str) -> Result<String> {
self.mgr.openapi_spec_of(gear_name).ok_or_else(|| {
DirectoryNotFound::new(format!("openapi spec for gear {gear_name}")).into()
})
}
async fn list_instances(&self, gear: &str) -> Result<Vec<ServiceInstanceInfo>> {
Ok(self
.mgr
.instances_of(gear)
.iter()
.map(|inst| project_instance(gear, inst))
.collect())
}
async fn resolve_by_labels(
&self,
gear: &str,
selector: &LabelSelector,
) -> Result<Vec<ServiceInstanceInfo>> {
Ok(self
.mgr
.instances_of(gear)
.iter()
.map(|inst| project_instance(gear, inst))
.filter(|i| selector.matches(&i.labels))
.collect())
}
async fn list_all_instances(&self) -> Result<Vec<ServiceInstanceInfo>> {
Ok(self
.mgr
.all_instances()
.iter()
.filter_map(|inst| {
let info = project_instance(&inst.gear, inst);
if info.endpoint.is_none() {
tracing::debug!(
gear = %inst.gear,
instance_id = %inst.instance_id,
"skipping instance with no gRPC or REST endpoint from cross-gear snapshot"
);
return None;
}
Some(info.without_labels())
})
.collect())
}
async fn register_instance(&self, info: RegisterInstanceInfo) -> Result<()> {
let instance_id = Uuid::parse_str(&info.instance_id).map_err(|e| {
anyhow::Error::from(DirectoryInvalidArgument::new(format!(
"invalid instance_id '{}': {e}",
info.instance_id
)))
})?;
cf_system_sdks::directory::validate_labels(&info.labels)
.map_err(|e| anyhow::Error::from(DirectoryInvalidArgument::new(e.to_string())))?;
let mut instance = GearInstance::new(info.gear.clone(), instance_id);
if let Some(version) = info.version {
instance = instance.with_version(version);
}
for (service_name, endpoint) in info.grpc_services {
instance = instance.with_grpc_service(service_name, Endpoint::from_uri(endpoint.uri));
}
if let Some(rest) = info.rest_endpoint {
instance = instance.with_rest_endpoint(Endpoint::from_uri(rest.uri));
}
if let Some(spec) = info.openapi_spec {
instance = instance.with_openapi_spec(spec);
}
if !info.labels.is_empty() {
instance = instance.with_labels(info.labels);
}
self.mgr.register_instance(Arc::new(instance)).map_err(
|GrpcServiceNameConflict {
service_name,
owner,
recoverable,
}| {
anyhow::Error::from(DirectoryServiceNameConflict {
service_name,
owner,
recoverable,
})
},
)?;
Ok(())
}
async fn deregister_instance(&self, gear: &str, instance_id: &str) -> Result<()> {
let instance_id = Uuid::parse_str(instance_id)
.map_err(|e| anyhow::anyhow!("Invalid instance_id '{instance_id}': {e}"))?;
self.mgr.deregister(gear, instance_id);
Ok(())
}
async fn send_heartbeat(&self, gear: &str, instance_id: &str) -> Result<()> {
let instance_id = Uuid::parse_str(instance_id)
.map_err(|e| anyhow::anyhow!("Invalid instance_id '{instance_id}': {e}"))?;
self.mgr
.update_heartbeat(gear, instance_id, std::time::Instant::now());
Ok(())
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests {
use super::*;
#[tokio::test]
async fn test_resolve_grpc_service_not_found() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(dir);
let err = api
.resolve_grpc_service("nonexistent.Service")
.await
.unwrap_err();
assert!(
err.downcast_ref::<DirectoryNotFound>().is_some(),
"expected the typed not-found sentinel, got: {err:?}"
);
}
#[tokio::test]
async fn test_register_instance_via_api() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(dir.clone());
let instance_id = Uuid::new_v4();
let register_info = RegisterInstanceInfo::new("test_gear", instance_id.to_string())
.with_grpc_services(vec![(
"test.Service".to_owned(),
ServiceEndpoint::http("127.0.0.1", 8001),
)])
.with_version("1.0.0");
api.register_instance(register_info).await.unwrap();
let instances = dir.instances_of("test_gear");
assert_eq!(instances.len(), 1);
assert_eq!(instances[0].instance_id, instance_id);
assert_eq!(instances[0].version, Some("1.0.0".to_owned()));
assert!(instances[0].grpc_services.contains_key("test.Service"));
}
#[tokio::test]
async fn register_rejects_invalid_labels_at_store_boundary() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
let err = api
.register_instance(
RegisterInstanceInfo::new("worker", Uuid::new_v4().to_string())
.with_rest_endpoint(ServiceEndpoint::new("http://worker:8080"))
.with_labels(labels(&[("bad key", "7")])),
)
.await
.expect_err("an invalid label must be rejected at the store boundary");
assert!(
err.downcast_ref::<DirectoryInvalidArgument>().is_some(),
"store-boundary label rejection must be typed InvalidArgument, got: {err}"
);
assert!(
dir.instances_of("worker").is_empty(),
"a rejected registration must not be stored"
);
}
#[tokio::test]
async fn test_register_and_resolve_rest_and_openapi() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(dir.clone());
let instance_id = Uuid::new_v4();
let register_info = RegisterInstanceInfo::new("billing", instance_id.to_string())
.with_version("1.0.0")
.with_rest_endpoint(ServiceEndpoint::http("billing", 8080))
.with_openapi_spec("{\"openapi\":\"3.1.0\"}");
api.register_instance(register_info).await.unwrap();
let resolved = api.resolve_rest_service("billing").await.unwrap();
assert_eq!(resolved.uri, concat!("http", "://billing:8080"));
let spec = api.get_openapi_spec("billing").await.unwrap();
assert!(spec.contains("openapi"));
}
#[tokio::test]
async fn test_resolve_rest_and_openapi_not_found() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(dir);
let rest_err = api.resolve_rest_service("missing").await.unwrap_err();
assert!(
rest_err.downcast_ref::<DirectoryNotFound>().is_some(),
"expected the typed not-found sentinel, got: {rest_err:?}"
);
let spec_err = api.get_openapi_spec("missing").await.unwrap_err();
assert!(
spec_err.downcast_ref::<DirectoryNotFound>().is_some(),
"expected the typed not-found sentinel, got: {spec_err:?}"
);
}
#[tokio::test]
async fn test_deregister_instance_via_api() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(dir.clone());
let instance_id = Uuid::new_v4();
let inst = Arc::new(GearInstance::new("test_gear", instance_id));
dir.register_instance(inst).unwrap();
assert_eq!(dir.instances_of("test_gear").len(), 1);
api.deregister_instance("test_gear", &instance_id.to_string())
.await
.unwrap();
assert_eq!(dir.instances_of("test_gear").len(), 0);
}
#[tokio::test]
async fn test_send_heartbeat_via_api() {
use crate::runtime::InstanceState;
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(dir.clone());
let instance_id = Uuid::new_v4();
let inst = Arc::new(GearInstance::new("test_gear", instance_id));
dir.register_instance(inst).unwrap();
let instances = dir.instances_of("test_gear");
assert_eq!(instances[0].state(), InstanceState::Registered);
api.send_heartbeat("test_gear", &instance_id.to_string())
.await
.unwrap();
let instances = dir.instances_of("test_gear");
assert_eq!(instances[0].state(), InstanceState::Healthy);
}
#[tokio::test]
async fn test_list_all_instances_across_gears() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
for (gear, port) in [("billing", 8080u16), ("catalog", 8081u16)] {
api.register_instance(
RegisterInstanceInfo::new(gear, Uuid::new_v4().to_string())
.with_version("1.0.0")
.with_rest_endpoint(ServiceEndpoint::http(gear, port))
.with_openapi_spec(format!("{{\"openapi\":\"3.1.0\",\"x\":\"{gear}\"}}")),
)
.await
.unwrap();
}
api.register_instance(
RegisterInstanceInfo::new("reporting", Uuid::new_v4().to_string())
.with_grpc_services(vec![(
"reporting.Service".to_owned(),
ServiceEndpoint::new("http://reporting:7000"),
)])
.with_version("1.0.0"),
)
.await
.unwrap();
let all = api.list_all_instances().await.unwrap();
assert_eq!(all.len(), 3);
let billing = all.iter().find(|i| i.gear == "billing").expect("billing");
assert_eq!(
billing.rest_endpoint.as_ref().map(|e| e.uri.as_str()),
Some("http://billing:8080")
);
assert!(billing.openapi_spec_hash.is_some());
assert!(
api.get_openapi_spec("billing")
.await
.expect("billing spec")
.contains("billing")
);
let reporting = all
.iter()
.find(|i| i.gear == "reporting")
.expect("reporting");
assert_eq!(
reporting.endpoint.as_ref().map(|e| e.uri.as_str()),
Some("http://reporting:7000")
);
assert!(reporting.rest_endpoint.is_none());
assert!(reporting.openapi_spec_hash.is_none());
}
fn labels(pairs: &[(&str, &str)]) -> std::collections::BTreeMap<String, String> {
pairs
.iter()
.map(|(k, v)| ((*k).to_owned(), (*v).to_owned()))
.collect()
}
#[tokio::test]
async fn register_carries_labels_into_list_instances() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
api.register_instance(
RegisterInstanceInfo::new("worker", Uuid::new_v4().to_string())
.with_grpc_services(vec![(
"worker.Svc".to_owned(),
ServiceEndpoint::new("http://worker:7000"),
)])
.with_labels(labels(&[("shard", "7")])),
)
.await
.unwrap();
let listed = api.list_instances("worker").await.unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(listed[0].labels.get("shard"), Some(&"7".to_owned()));
}
#[tokio::test]
async fn resolve_by_labels_selects_matching_instances() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
for (id, shard) in [("a", "7"), ("b", "8"), ("c", "7")] {
api.register_instance(
RegisterInstanceInfo::new("worker", Uuid::new_v4().to_string())
.with_grpc_services(vec![(
format!("worker.{id}"),
ServiceEndpoint::new(format!("http://worker-{id}:7000")),
)])
.with_labels(labels(&[("shard", shard)])),
)
.await
.unwrap();
}
let matched = api
.resolve_by_labels("worker", &LabelSelector::new().with("shard", "7"))
.await
.unwrap();
assert_eq!(matched.len(), 2, "two instances carry shard=7");
assert!(
matched
.iter()
.all(|i| i.labels.get("shard") == Some(&"7".to_owned()))
);
}
#[tokio::test]
async fn labelless_reregister_preserves_labels_through_api() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
let instance_id = Uuid::new_v4().to_string();
api.register_instance(
RegisterInstanceInfo::new("worker", instance_id.clone())
.with_rest_endpoint(ServiceEndpoint::new("http://worker:8080"))
.with_labels(labels(&[("shard", "7")])),
)
.await
.unwrap();
api.register_instance(
RegisterInstanceInfo::new("worker", instance_id)
.with_rest_endpoint(ServiceEndpoint::new("http://worker:8080"))
.with_version("2.0.0"),
)
.await
.unwrap();
let listed = api.list_instances("worker").await.unwrap();
assert_eq!(listed.len(), 1);
assert_eq!(
listed[0].labels.get("shard"),
Some(&"7".to_owned()),
"label-less re-registration must not wipe stored labels"
);
let matched = api
.resolve_by_labels("worker", &LabelSelector::new().with("shard", "7"))
.await
.unwrap();
assert_eq!(matched.len(), 1);
}
#[tokio::test]
async fn resolve_by_labels_carries_live_serving_state() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
let instance_id = Uuid::new_v4();
api.register_instance(
RegisterInstanceInfo::new("worker", instance_id.to_string())
.with_grpc_services(vec![(
"worker.Svc".to_owned(),
ServiceEndpoint::new("http://worker:7000"),
)])
.with_labels(labels(&[("shard", "7")])),
)
.await
.unwrap();
let matched = api
.resolve_by_labels("worker", &LabelSelector::new().with("shard", "7"))
.await
.unwrap();
assert_eq!(matched.len(), 1);
assert_eq!(matched[0].state, InstanceState::Registered);
assert!(!matched[0].state.is_serving());
dir.update_heartbeat("worker", instance_id, std::time::Instant::now());
let matched = api
.resolve_by_labels("worker", &LabelSelector::new().with("shard", "7"))
.await
.unwrap();
assert_eq!(matched[0].state, InstanceState::Healthy);
assert!(matched[0].state.is_serving());
}
#[tokio::test]
async fn resolve_by_labels_omits_openapi_spec() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
api.register_instance(
RegisterInstanceInfo::new("worker", Uuid::new_v4().to_string())
.with_rest_endpoint(ServiceEndpoint::new("http://worker:8080"))
.with_openapi_spec("{\"openapi\":\"3.1.0\"}")
.with_labels(labels(&[("shard", "7")])),
)
.await
.unwrap();
let matched = api
.resolve_by_labels("worker", &LabelSelector::new().with("shard", "7"))
.await
.unwrap();
assert_eq!(matched.len(), 1);
assert!(
matched[0].openapi_spec_hash.is_some(),
"in-process resolve_by_labels carries only the spec hash, never the document"
);
let listed = api.list_instances("worker").await.unwrap();
assert!(
listed[0].openapi_spec_hash.is_some(),
"list_instances must carry the spec hash, never the document"
);
}
#[tokio::test]
async fn list_all_instances_skips_endpoint_less() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
api.register_instance(
RegisterInstanceInfo::new("worker", Uuid::new_v4().to_string())
.with_rest_endpoint(ServiceEndpoint::new("http://worker:8080")),
)
.await
.unwrap();
api.register_instance(RegisterInstanceInfo::new(
"placeholder",
Uuid::new_v4().to_string(),
))
.await
.unwrap();
let all = api.list_all_instances().await.unwrap();
assert_eq!(all.len(), 1, "the endpoint-less instance is skipped");
assert_eq!(all[0].gear, "worker");
assert!(
all.iter()
.all(|i| i.endpoint.as_ref().is_some_and(|e| !e.uri.is_empty())),
"no instance may carry an absent or empty-URI endpoint"
);
}
#[tokio::test]
async fn label_path_returns_endpoint_less_match() {
let dir = Arc::new(GearManager::new());
let api = LocalDirectoryClient::new(Arc::clone(&dir));
api.register_instance(
RegisterInstanceInfo::new("worker", Uuid::new_v4().to_string())
.with_labels(labels(&[("shard", "7")])),
)
.await
.unwrap();
let matched = api
.resolve_by_labels("worker", &LabelSelector::new().with("shard", "7"))
.await
.unwrap();
assert_eq!(
matched.len(),
1,
"resolve_by_labels must return a matched instance even with no endpoint"
);
assert!(
matched[0].endpoint.is_none(),
"a no-endpoint match carries endpoint = None, not an empty-URI sentinel"
);
assert!(matched[0].rest_endpoint.is_none());
assert_eq!(matched[0].labels.get("shard"), Some(&"7".to_owned()));
let listed = api.list_instances("worker").await.unwrap();
assert_eq!(
listed.len(),
1,
"list_instances must not drop endpoint-less instances"
);
assert!(listed[0].endpoint.is_none());
}
}