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};
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, GrpcServiceInfo,
RegisterInstanceInfo, ServiceEndpoint, ServiceInstanceInfo,
};
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>> {
let mut result = Vec::new();
for inst in self.mgr.instances_of(gear) {
if let Some((_, ep)) = inst.grpc_services.iter().next() {
result.push(ServiceInstanceInfo {
gear: gear.to_owned(),
instance_id: inst.instance_id.to_string(),
endpoint: ServiceEndpoint::new(ep.uri.clone()),
version: inst.version.clone(),
rest_endpoint: inst
.rest_endpoint
.as_ref()
.map(|ep| ServiceEndpoint::new(ep.uri.clone())),
openapi_spec_hash: inst.openapi_spec.as_deref().map(openapi_spec_hash),
openapi_spec: inst.openapi_spec.clone(),
grpc_services: inst
.grpc_services
.iter()
.map(|(name, e)| (name.clone(), ServiceEndpoint::new(e.uri.clone())))
.collect(),
});
}
}
Ok(result)
}
async fn list_all_instances(&self) -> Result<Vec<ServiceInstanceInfo>> {
let result = self
.mgr
.all_instances()
.into_iter()
.map(|inst| {
let endpoint = inst
.grpc_services
.values()
.next()
.or(inst.rest_endpoint.as_ref())
.map_or_else(
|| ServiceEndpoint::new(String::new()),
|ep| ServiceEndpoint::new(ep.uri.clone()),
);
ServiceInstanceInfo {
gear: inst.gear.clone(),
instance_id: inst.instance_id.to_string(),
endpoint,
version: inst.version.clone(),
rest_endpoint: inst
.rest_endpoint
.as_ref()
.map(|ep| ServiceEndpoint::new(ep.uri.clone())),
openapi_spec_hash: inst.openapi_spec.as_deref().map(openapi_spec_hash),
openapi_spec: None,
grpc_services: inst
.grpc_services
.iter()
.map(|(name, e)| (name.clone(), ServiceEndpoint::new(e.uri.clone())))
.collect(),
}
})
.collect();
Ok(result)
}
async fn register_instance(&self, info: RegisterInstanceInfo) -> Result<()> {
let instance_id = Uuid::parse_str(&info.instance_id)
.map_err(|e| anyhow::anyhow!("Invalid instance_id '{}': {}", info.instance_id, e))?;
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);
}
self.mgr.register_instance(Arc::new(instance));
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 {
gear: "test_gear".to_owned(),
instance_id: instance_id.to_string(),
grpc_services: vec![(
"test.Service".to_owned(),
ServiceEndpoint::http("127.0.0.1", 8001),
)],
version: Some("1.0.0".to_owned()),
rest_endpoint: None,
openapi_spec: None,
};
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 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 {
gear: "billing".to_owned(),
instance_id: instance_id.to_string(),
grpc_services: vec![],
version: Some("1.0.0".to_owned()),
rest_endpoint: Some(ServiceEndpoint::http("billing", 8080)),
openapi_spec: Some("{\"openapi\":\"3.1.0\"}".to_owned()),
};
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);
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);
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 {
gear: gear.to_owned(),
instance_id: Uuid::new_v4().to_string(),
grpc_services: vec![],
version: Some("1.0.0".to_owned()),
rest_endpoint: Some(ServiceEndpoint::http(gear, port)),
openapi_spec: Some(format!("{{\"openapi\":\"3.1.0\",\"x\":\"{gear}\"}}")),
})
.await
.unwrap();
}
api.register_instance(RegisterInstanceInfo {
gear: "reporting".to_owned(),
instance_id: Uuid::new_v4().to_string(),
grpc_services: vec![(
"reporting.Service".to_owned(),
ServiceEndpoint::new("http://reporting:7000"),
)],
version: Some("1.0.0".to_owned()),
rest_endpoint: None,
openapi_spec: None,
})
.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.is_none());
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.uri.as_str(), "http://reporting:7000");
assert!(reporting.rest_endpoint.is_none());
assert!(reporting.openapi_spec.is_none());
}
}