1use anyhow::Result;
4use async_trait::async_trait;
5use std::hash::{Hash, Hasher};
6use std::sync::Arc;
7use uuid::Uuid;
8
9use crate::runtime::{Endpoint, GearInstance, GearManager};
10
11fn openapi_spec_hash(spec: &str) -> String {
28 let mut hasher = std::collections::hash_map::DefaultHasher::new();
29 spec.hash(&mut hasher);
30 format!("{:016x}", hasher.finish())
31}
32
33pub use cf_system_sdks::directory::{
35 DirectoryClient, DirectoryInvalidArgument, DirectoryNotFound, GrpcServiceInfo,
36 RegisterInstanceInfo, ServiceEndpoint, ServiceInstanceInfo,
37};
38
39pub struct LocalDirectoryClient {
44 mgr: Arc<GearManager>,
45}
46
47impl LocalDirectoryClient {
48 #[must_use]
49 pub fn new(mgr: Arc<GearManager>) -> Self {
50 Self { mgr }
51 }
52}
53
54#[async_trait]
55impl DirectoryClient for LocalDirectoryClient {
56 async fn resolve_grpc_service(&self, service_name: &str) -> Result<ServiceEndpoint> {
64 if let Some((_gear, _inst, ep)) = self.mgr.pick_service_round_robin(service_name) {
65 return Ok(ServiceEndpoint::new(ep.uri));
66 }
67
68 Err(DirectoryNotFound::new(format!("service {service_name}")).into())
69 }
70
71 async fn resolve_rest_service(&self, gear_name: &str) -> Result<ServiceEndpoint> {
72 if let Some(ep) = self.mgr.pick_rest_endpoint_round_robin(gear_name) {
73 return Ok(ServiceEndpoint::new(ep.uri));
74 }
75
76 Err(DirectoryNotFound::new(format!("gear {gear_name}")).into())
77 }
78
79 async fn get_openapi_spec(&self, gear_name: &str) -> Result<String> {
80 self.mgr.openapi_spec_of(gear_name).ok_or_else(|| {
81 DirectoryNotFound::new(format!("openapi spec for gear {gear_name}")).into()
82 })
83 }
84
85 async fn list_instances(&self, gear: &str) -> Result<Vec<ServiceInstanceInfo>> {
86 let mut result = Vec::new();
87
88 for inst in self.mgr.instances_of(gear) {
89 if let Some((_, ep)) = inst.grpc_services.iter().next() {
90 result.push(ServiceInstanceInfo {
91 gear: gear.to_owned(),
92 instance_id: inst.instance_id.to_string(),
93 endpoint: ServiceEndpoint::new(ep.uri.clone()),
94 version: inst.version.clone(),
95 rest_endpoint: inst
96 .rest_endpoint
97 .as_ref()
98 .map(|ep| ServiceEndpoint::new(ep.uri.clone())),
99 openapi_spec_hash: inst.openapi_spec.as_deref().map(openapi_spec_hash),
100 openapi_spec: inst.openapi_spec.clone(),
101 grpc_services: inst
105 .grpc_services
106 .iter()
107 .map(|(name, e)| (name.clone(), ServiceEndpoint::new(e.uri.clone())))
108 .collect(),
109 });
110 }
111 }
112
113 Ok(result)
114 }
115
116 async fn list_all_instances(&self) -> Result<Vec<ServiceInstanceInfo>> {
117 let result = self
118 .mgr
119 .all_instances()
120 .into_iter()
121 .map(|inst| {
122 let endpoint = inst
125 .grpc_services
126 .values()
127 .next()
128 .or(inst.rest_endpoint.as_ref())
129 .map_or_else(
130 || ServiceEndpoint::new(String::new()),
131 |ep| ServiceEndpoint::new(ep.uri.clone()),
132 );
133 ServiceInstanceInfo {
134 gear: inst.gear.clone(),
135 instance_id: inst.instance_id.to_string(),
136 endpoint,
137 version: inst.version.clone(),
138 rest_endpoint: inst
139 .rest_endpoint
140 .as_ref()
141 .map(|ep| ServiceEndpoint::new(ep.uri.clone())),
142 openapi_spec_hash: inst.openapi_spec.as_deref().map(openapi_spec_hash),
149 openapi_spec: None,
150 grpc_services: inst
155 .grpc_services
156 .iter()
157 .map(|(name, e)| (name.clone(), ServiceEndpoint::new(e.uri.clone())))
158 .collect(),
159 }
160 })
161 .collect();
162
163 Ok(result)
164 }
165
166 async fn register_instance(&self, info: RegisterInstanceInfo) -> Result<()> {
167 let instance_id = Uuid::parse_str(&info.instance_id)
169 .map_err(|e| anyhow::anyhow!("Invalid instance_id '{}': {}", info.instance_id, e))?;
170
171 let mut instance = GearInstance::new(info.gear.clone(), instance_id);
173
174 if let Some(version) = info.version {
176 instance = instance.with_version(version);
177 }
178
179 for (service_name, endpoint) in info.grpc_services {
181 instance = instance.with_grpc_service(service_name, Endpoint::from_uri(endpoint.uri));
182 }
183
184 if let Some(rest) = info.rest_endpoint {
186 instance = instance.with_rest_endpoint(Endpoint::from_uri(rest.uri));
187 }
188
189 if let Some(spec) = info.openapi_spec {
191 instance = instance.with_openapi_spec(spec);
192 }
193
194 self.mgr.register_instance(Arc::new(instance));
196
197 Ok(())
198 }
199
200 async fn deregister_instance(&self, gear: &str, instance_id: &str) -> Result<()> {
201 let instance_id = Uuid::parse_str(instance_id)
202 .map_err(|e| anyhow::anyhow!("Invalid instance_id '{instance_id}': {e}"))?;
203 self.mgr.deregister(gear, instance_id);
204 Ok(())
205 }
206
207 async fn send_heartbeat(&self, gear: &str, instance_id: &str) -> Result<()> {
208 let instance_id = Uuid::parse_str(instance_id)
209 .map_err(|e| anyhow::anyhow!("Invalid instance_id '{instance_id}': {e}"))?;
210 self.mgr
211 .update_heartbeat(gear, instance_id, std::time::Instant::now());
212 Ok(())
213 }
214}
215
216#[cfg(test)]
217#[cfg_attr(coverage_nightly, coverage(off))]
218mod tests {
219 use super::*;
220
221 #[tokio::test]
222 async fn test_resolve_grpc_service_not_found() {
223 let dir = Arc::new(GearManager::new());
224 let api = LocalDirectoryClient::new(dir);
225
226 let err = api
227 .resolve_grpc_service("nonexistent.Service")
228 .await
229 .unwrap_err();
230 assert!(
234 err.downcast_ref::<DirectoryNotFound>().is_some(),
235 "expected the typed not-found sentinel, got: {err:?}"
236 );
237 }
238
239 #[tokio::test]
240 async fn test_register_instance_via_api() {
241 let dir = Arc::new(GearManager::new());
242 let api = LocalDirectoryClient::new(dir.clone());
243
244 let instance_id = Uuid::new_v4();
245 let register_info = RegisterInstanceInfo {
247 gear: "test_gear".to_owned(),
248 instance_id: instance_id.to_string(),
249 grpc_services: vec![(
250 "test.Service".to_owned(),
251 ServiceEndpoint::http("127.0.0.1", 8001),
252 )],
253 version: Some("1.0.0".to_owned()),
254 rest_endpoint: None,
255 openapi_spec: None,
256 };
257
258 api.register_instance(register_info).await.unwrap();
259
260 let instances = dir.instances_of("test_gear");
262 assert_eq!(instances.len(), 1);
263 assert_eq!(instances[0].instance_id, instance_id);
264 assert_eq!(instances[0].version, Some("1.0.0".to_owned()));
265 assert!(instances[0].grpc_services.contains_key("test.Service"));
266 }
267
268 #[tokio::test]
269 async fn test_register_and_resolve_rest_and_openapi() {
270 let dir = Arc::new(GearManager::new());
271 let api = LocalDirectoryClient::new(dir.clone());
272
273 let instance_id = Uuid::new_v4();
274 let register_info = RegisterInstanceInfo {
275 gear: "billing".to_owned(),
276 instance_id: instance_id.to_string(),
277 grpc_services: vec![],
278 version: Some("1.0.0".to_owned()),
279 rest_endpoint: Some(ServiceEndpoint::http("billing", 8080)),
280 openapi_spec: Some("{\"openapi\":\"3.1.0\"}".to_owned()),
281 };
282
283 api.register_instance(register_info).await.unwrap();
284
285 let resolved = api.resolve_rest_service("billing").await.unwrap();
287 assert_eq!(resolved.uri, concat!("http", "://billing:8080"));
288
289 let spec = api.get_openapi_spec("billing").await.unwrap();
291 assert!(spec.contains("openapi"));
292 }
293
294 #[tokio::test]
295 async fn test_resolve_rest_and_openapi_not_found() {
296 let dir = Arc::new(GearManager::new());
297 let api = LocalDirectoryClient::new(dir);
298
299 let rest_err = api.resolve_rest_service("missing").await.unwrap_err();
302 assert!(
303 rest_err.downcast_ref::<DirectoryNotFound>().is_some(),
304 "expected the typed not-found sentinel, got: {rest_err:?}"
305 );
306
307 let spec_err = api.get_openapi_spec("missing").await.unwrap_err();
308 assert!(
309 spec_err.downcast_ref::<DirectoryNotFound>().is_some(),
310 "expected the typed not-found sentinel, got: {spec_err:?}"
311 );
312 }
313
314 #[tokio::test]
315 async fn test_deregister_instance_via_api() {
316 let dir = Arc::new(GearManager::new());
317 let api = LocalDirectoryClient::new(dir.clone());
318
319 let instance_id = Uuid::new_v4();
320 let inst = Arc::new(GearInstance::new("test_gear", instance_id));
322 dir.register_instance(inst);
323
324 assert_eq!(dir.instances_of("test_gear").len(), 1);
326
327 api.deregister_instance("test_gear", &instance_id.to_string())
329 .await
330 .unwrap();
331
332 assert_eq!(dir.instances_of("test_gear").len(), 0);
334 }
335
336 #[tokio::test]
337 async fn test_send_heartbeat_via_api() {
338 use crate::runtime::InstanceState;
339
340 let dir = Arc::new(GearManager::new());
341 let api = LocalDirectoryClient::new(dir.clone());
342
343 let instance_id = Uuid::new_v4();
344 let inst = Arc::new(GearInstance::new("test_gear", instance_id));
346 dir.register_instance(inst);
347
348 let instances = dir.instances_of("test_gear");
350 assert_eq!(instances[0].state(), InstanceState::Registered);
351
352 api.send_heartbeat("test_gear", &instance_id.to_string())
354 .await
355 .unwrap();
356
357 let instances = dir.instances_of("test_gear");
359 assert_eq!(instances[0].state(), InstanceState::Healthy);
360 }
361
362 #[tokio::test]
363 async fn test_list_all_instances_across_gears() {
364 let dir = Arc::new(GearManager::new());
365 let api = LocalDirectoryClient::new(Arc::clone(&dir));
366
367 for (gear, port) in [("billing", 8080u16), ("catalog", 8081u16)] {
369 api.register_instance(RegisterInstanceInfo {
370 gear: gear.to_owned(),
371 instance_id: Uuid::new_v4().to_string(),
372 grpc_services: vec![],
373 version: Some("1.0.0".to_owned()),
374 rest_endpoint: Some(ServiceEndpoint::http(gear, port)),
375 openapi_spec: Some(format!("{{\"openapi\":\"3.1.0\",\"x\":\"{gear}\"}}")),
376 })
377 .await
378 .unwrap();
379 }
380
381 api.register_instance(RegisterInstanceInfo {
385 gear: "reporting".to_owned(),
386 instance_id: Uuid::new_v4().to_string(),
387 grpc_services: vec![(
388 "reporting.Service".to_owned(),
389 ServiceEndpoint::new("http://reporting:7000"),
390 )],
391 version: Some("1.0.0".to_owned()),
392 rest_endpoint: None,
393 openapi_spec: None,
394 })
395 .await
396 .unwrap();
397
398 let all = api.list_all_instances().await.unwrap();
399 assert_eq!(all.len(), 3);
400
401 let billing = all.iter().find(|i| i.gear == "billing").expect("billing");
402 assert_eq!(
403 billing.rest_endpoint.as_ref().map(|e| e.uri.as_str()),
404 Some("http://billing:8080")
405 );
406 assert!(billing.openapi_spec.is_none());
409 assert!(
410 api.get_openapi_spec("billing")
411 .await
412 .expect("billing spec")
413 .contains("billing")
414 );
415
416 let reporting = all
419 .iter()
420 .find(|i| i.gear == "reporting")
421 .expect("reporting");
422 assert_eq!(reporting.endpoint.uri.as_str(), "http://reporting:7000");
423 assert!(reporting.rest_endpoint.is_none());
424 assert!(reporting.openapi_spec.is_none());
425 }
426}