1use anyhow::Result;
6use async_trait::async_trait;
7use tonic::service::interceptor::InterceptedService;
8use tonic::transport::Channel;
9
10use crate::ProtoInstanceState;
11use crate::api::{
12 DirectoryClient, DirectoryInvalidArgument, DirectoryNotFound, DirectoryPermissionDenied,
13 InstanceState, LabelSelector, RegisterInstanceInfo, ServiceEndpoint, ServiceInstanceInfo,
14};
15use std::collections::BTreeMap;
16use toolkit_transport_grpc::InternalAuthInterceptor;
17use toolkit_transport_grpc::client::{GrpcClientConfig, connect_lazy, connect_with_retry};
18
19use crate::{
20 DeregisterInstanceRequest, DirectoryServiceClient, GetOpenApiSpecRequest, GrpcServiceEndpoint,
21 HeartbeatRequest, InstanceInfo, ListAllInstancesRequest, ListInstancesRequest,
22 RegisterInstanceRequest, ResolveGrpcServiceRequest, ResolveRestServiceRequest,
23};
24
25type AuthedChannel = InterceptedService<Channel, InternalAuthInterceptor>;
31
32fn lookup_error(resource: &str, status: &tonic::Status) -> anyhow::Error {
40 match status.code() {
41 tonic::Code::NotFound => DirectoryNotFound::new(resource.to_owned()).into(),
42 tonic::Code::InvalidArgument => {
43 DirectoryInvalidArgument::new(status.message().to_owned()).into()
44 }
45 code => anyhow::anyhow!(
46 "directory lookup for {resource} failed: gRPC {code:?}: {}",
47 status.message()
48 ),
49 }
50}
51
52fn call_error(op: &str, status: &tonic::Status) -> anyhow::Error {
72 match status.code() {
73 tonic::Code::InvalidArgument => {
74 DirectoryInvalidArgument::new(status.message().to_owned()).into()
75 }
76 tonic::Code::PermissionDenied => {
77 DirectoryPermissionDenied::new(status.message().to_owned()).into()
78 }
79 code => anyhow::anyhow!("directory {op} failed: gRPC {code:?}: {}", status.message()),
80 }
81}
82
83pub struct DirectoryGrpcClient {
94 inner: DirectoryServiceClient<AuthedChannel>,
95}
96
97impl DirectoryGrpcClient {
98 pub async fn connect(uri: impl Into<String>) -> Result<Self> {
106 let cfg = GrpcClientConfig::new("directory");
107 Self::connect_with_retry(uri, &cfg).await
108 }
109
110 pub async fn connect_with_interceptor(
122 uri: impl Into<String>,
123 interceptor: InternalAuthInterceptor,
124 ) -> Result<Self> {
125 let cfg = GrpcClientConfig::new("directory");
126 let channel: Channel = connect_with_retry(uri, &cfg).await?;
127 Ok(Self::from_channel_with_interceptor(channel, interceptor))
128 }
129
130 pub fn connect_lazy(uri: impl Into<String>) -> Result<Self> {
149 let cfg = GrpcClientConfig::new("directory");
150 let channel: Channel = connect_lazy(uri, &cfg)?;
151 Ok(Self::from_channel(channel))
152 }
153
154 pub fn connect_lazy_with_interceptor(
166 uri: impl Into<String>,
167 interceptor: InternalAuthInterceptor,
168 ) -> Result<Self> {
169 let cfg = GrpcClientConfig::new("directory");
170 let channel: Channel = connect_lazy(uri, &cfg)?;
172 Ok(Self::from_channel_with_interceptor(channel, interceptor))
173 }
174
175 pub async fn connect_with_retry(
183 uri: impl Into<String>,
184 cfg: &GrpcClientConfig,
185 ) -> Result<Self> {
186 let channel: Channel = connect_with_retry(uri, cfg).await?;
187 Ok(Self::from_channel(channel))
188 }
189
190 pub async fn connect_no_retry(uri: impl Into<String>, cfg: &GrpcClientConfig) -> Result<Self> {
198 let uri_string = uri.into();
199
200 let endpoint = tonic::transport::Endpoint::from_shared(uri_string)?
202 .connect_timeout(cfg.connect_timeout)
203 .timeout(cfg.rpc_timeout);
204
205 let channel = endpoint.connect().await?;
207
208 if cfg.enable_tracing {
209 tracing::debug!(
210 service_name = cfg.service_name,
211 connect_timeout_ms = cfg.connect_timeout.as_millis(),
212 rpc_timeout_ms = cfg.rpc_timeout.as_millis(),
213 "directory gRPC client connected"
214 );
215 }
216
217 Ok(Self::from_channel(channel))
218 }
219
220 #[must_use]
226 pub fn from_channel(channel: Channel) -> Self {
227 Self::from_channel_with_interceptor(channel, InternalAuthInterceptor::disabled())
228 }
229
230 #[must_use]
233 pub fn from_channel_with_interceptor(
234 channel: Channel,
235 interceptor: InternalAuthInterceptor,
236 ) -> Self {
237 Self {
238 inner: DirectoryServiceClient::with_interceptor(channel, interceptor),
239 }
240 }
241}
242
243#[async_trait]
244impl DirectoryClient for DirectoryGrpcClient {
245 async fn resolve_grpc_service(&self, service_name: &str) -> Result<ServiceEndpoint> {
246 let mut client = self.inner.clone();
247 let request = tonic::Request::new(ResolveGrpcServiceRequest {
248 service_name: service_name.to_owned(),
249 });
250
251 let response = client
252 .resolve_grpc_service(request)
253 .await
254 .map_err(|e| lookup_error(&format!("service {service_name}"), &e))?;
255
256 let proto_response = response.into_inner();
257 Ok(ServiceEndpoint::new(proto_response.endpoint_uri))
258 }
259
260 async fn resolve_rest_service(&self, gear_name: &str) -> Result<ServiceEndpoint> {
261 let mut client = self.inner.clone();
262 let request = tonic::Request::new(ResolveRestServiceRequest {
263 gear_name: gear_name.to_owned(),
264 });
265
266 let response = client
267 .resolve_rest_service(request)
268 .await
269 .map_err(|e| lookup_error(&format!("gear {gear_name}"), &e))?;
270
271 let proto_response = response.into_inner();
272 Ok(ServiceEndpoint::new(proto_response.endpoint_uri))
273 }
274
275 async fn get_openapi_spec(&self, gear_name: &str) -> Result<String> {
276 let mut client = self.inner.clone();
277 let request = tonic::Request::new(GetOpenApiSpecRequest {
278 gear_name: gear_name.to_owned(),
279 });
280
281 let response = client
282 .get_open_api_spec(request)
283 .await
284 .map_err(|e| lookup_error(&format!("openapi spec for gear {gear_name}"), &e))?;
285
286 Ok(response.into_inner().openapi_spec)
287 }
288
289 async fn list_instances(&self, gear: &str) -> Result<Vec<ServiceInstanceInfo>> {
290 let mut client = self.inner.clone();
291 let request = tonic::Request::new(ListInstancesRequest {
292 gear_name: gear.to_owned(),
293 match_labels: std::collections::HashMap::new(),
294 });
295
296 let response = client
297 .list_instances(request)
298 .await
299 .map_err(|e| lookup_error(&format!("instances of gear {gear}"), &e))?;
300
301 let instances = response
302 .into_inner()
303 .instances
304 .into_iter()
305 .map(proto_instance_to_domain)
306 .collect();
307
308 Ok(instances)
309 }
310
311 async fn resolve_by_labels(
312 &self,
313 gear: &str,
314 selector: &LabelSelector,
315 ) -> Result<Vec<ServiceInstanceInfo>> {
316 let mut client = self.inner.clone();
326 let match_labels = selector
327 .match_labels
328 .iter()
329 .map(|(k, v)| (k.clone(), v.clone()))
330 .collect();
331 let request = tonic::Request::new(ListInstancesRequest {
332 gear_name: gear.to_owned(),
333 match_labels,
334 });
335
336 let response = client
337 .list_instances(request)
338 .await
339 .map_err(|e| lookup_error(&format!("instances of gear {gear}"), &e))?;
340
341 let instances = response
342 .into_inner()
343 .instances
344 .into_iter()
345 .map(proto_instance_to_domain)
346 .filter(|i| selector.matches(&i.labels))
347 .collect();
348
349 Ok(instances)
350 }
351
352 async fn list_all_instances(&self) -> Result<Vec<ServiceInstanceInfo>> {
353 let mut client = self.inner.clone();
354 let response = client
355 .list_all_instances(tonic::Request::new(ListAllInstancesRequest {}))
356 .await
357 .map_err(|e| lookup_error("all instances", &e))?;
358
359 let instances = response
360 .into_inner()
361 .instances
362 .into_iter()
363 .map(|proto| proto_instance_to_domain(proto).without_labels())
364 .collect();
365
366 Ok(instances)
367 }
368
369 async fn register_instance(&self, info: RegisterInstanceInfo) -> Result<()> {
370 let mut client = self.inner.clone();
371
372 let grpc_services = info
374 .grpc_services
375 .into_iter()
376 .map(|(name, ep)| GrpcServiceEndpoint {
377 service_name: name,
378 endpoint_uri: ep.uri,
379 })
380 .collect();
381
382 let req = RegisterInstanceRequest {
383 gear_name: info.gear,
384 instance_id: info.instance_id,
385 grpc_services,
386 version: info.version.unwrap_or_default(),
387 rest_endpoint_uri: info.rest_endpoint.map(|ep| ep.uri),
388 openapi_spec: info.openapi_spec,
389 labels: info.labels.into_iter().collect(),
390 };
391
392 client
393 .register_instance(tonic::Request::new(req))
394 .await
395 .map_err(|e| call_error("register_instance", &e))?;
396
397 Ok(())
398 }
399
400 async fn deregister_instance(&self, gear: &str, instance_id: &str) -> Result<()> {
401 let mut client = self.inner.clone();
402
403 let req = DeregisterInstanceRequest {
404 gear_name: gear.to_owned(),
405 instance_id: instance_id.to_owned(),
406 };
407
408 client
409 .deregister_instance(tonic::Request::new(req))
410 .await
411 .map_err(|e| call_error("deregister_instance", &e))?;
412
413 Ok(())
414 }
415
416 async fn send_heartbeat(&self, gear: &str, instance_id: &str) -> Result<()> {
417 let mut client = self.inner.clone();
418
419 let req = HeartbeatRequest {
420 gear_name: gear.to_owned(),
421 instance_id: instance_id.to_owned(),
422 };
423
424 client
425 .heartbeat(tonic::Request::new(req))
426 .await
427 .map_err(|e| call_error("heartbeat", &e))?;
428
429 Ok(())
430 }
431}
432
433fn proto_instance_to_domain(proto: InstanceInfo) -> ServiceInstanceInfo {
435 ServiceInstanceInfo {
436 gear: proto.gear_name,
437 instance_id: proto.instance_id,
438 endpoint: if proto.endpoint_uri.is_empty() {
439 None
440 } else {
441 Some(ServiceEndpoint::new(proto.endpoint_uri))
442 },
443 version: if proto.version.is_empty() {
444 None
445 } else {
446 Some(proto.version)
447 },
448 rest_endpoint: proto.rest_endpoint_uri.map(ServiceEndpoint::new),
449 openapi_spec_hash: proto.openapi_spec_hash,
450 grpc_services: Vec::new(),
459 labels: proto.labels.into_iter().collect::<BTreeMap<_, _>>(),
462 state: proto_state_to_domain(proto.state),
464 }
465}
466
467fn proto_state_to_domain(state: i32) -> InstanceState {
477 match ProtoInstanceState::try_from(state) {
478 Ok(ProtoInstanceState::Ready) => InstanceState::Ready,
479 Ok(ProtoInstanceState::Healthy) => InstanceState::Healthy,
480 Ok(ProtoInstanceState::Quarantined) => InstanceState::Quarantined,
481 Ok(ProtoInstanceState::Draining) => InstanceState::Draining,
482 Ok(ProtoInstanceState::Registered) => InstanceState::Registered,
483 Ok(ProtoInstanceState::Unspecified) => InstanceState::Unknown,
484 Err(_) => {
485 tracing::warn!(
486 raw_state = state,
487 "directory returned an unrecognised InstanceState discriminant; \
488 treating as Unknown (non-serving)"
489 );
490 InstanceState::Unknown
491 }
492 }
493}
494
495#[cfg(test)]
496#[cfg_attr(coverage_nightly, coverage(off))]
497mod tests {
498 use super::*;
499
500 #[tokio::test]
501 async fn test_grpc_client_can_be_constructed() {
502 let endpoint = tonic::transport::Endpoint::from_static("http://[::1]:50051");
504
505 let channel_result = endpoint.connect().await;
508
509 if let Ok(channel) = channel_result {
511 let _client = DirectoryGrpcClient::from_channel(channel);
512 }
513 }
514
515 #[tokio::test]
516 async fn from_channel_constructs_without_connecting() {
517 let channel = Channel::from_static("http://[::1]:50051").connect_lazy();
521 let _default = DirectoryGrpcClient::from_channel(channel.clone());
522 let _authed = DirectoryGrpcClient::from_channel_with_interceptor(
523 channel,
524 InternalAuthInterceptor::disabled(),
525 );
526 }
527
528 #[tokio::test]
529 async fn connect_lazy_succeeds_against_unreachable_peer() {
530 let plain = DirectoryGrpcClient::connect_lazy("http://127.0.0.1:1");
535 assert!(
536 plain.is_ok(),
537 "connect_lazy must not eagerly connect (unreachable peer -> Ok)"
538 );
539
540 let authed = DirectoryGrpcClient::connect_lazy_with_interceptor(
541 "http://127.0.0.1:1",
542 InternalAuthInterceptor::disabled(),
543 );
544 assert!(
545 authed.is_ok(),
546 "connect_lazy_with_interceptor must not eagerly connect (unreachable peer -> Ok)"
547 );
548 }
549
550 #[tokio::test]
551 async fn connect_lazy_rejects_malformed_uri() {
552 assert!(
555 DirectoryGrpcClient::connect_lazy(String::new()).is_err(),
556 "connect_lazy should fail on a malformed URI"
557 );
558 assert!(
559 DirectoryGrpcClient::connect_lazy_with_interceptor(
560 String::new(),
561 InternalAuthInterceptor::disabled(),
562 )
563 .is_err(),
564 "connect_lazy_with_interceptor should fail on a malformed URI"
565 );
566 }
567
568 #[tokio::test]
569 async fn resolve_grpc_service_through_lazy_client_errors_not_hangs() {
570 let client =
574 DirectoryGrpcClient::connect_lazy("http://127.0.0.1:1").expect("lazy build ok");
575 let outcome = tokio::time::timeout(
576 std::time::Duration::from_secs(5),
577 client.resolve_grpc_service("cf.directory.v1.DirectoryService"),
578 )
579 .await;
580 assert!(
581 outcome.is_ok(),
582 "resolve_grpc_service through a lazy client must not hang against an unreachable peer"
583 );
584 assert!(
585 outcome.unwrap().is_err(),
586 "resolve_grpc_service against an unreachable directory must return Err"
587 );
588 }
589
590 #[test]
591 fn proto_instance_maps_all_fields_to_domain() {
592 let proto = InstanceInfo {
593 gear_name: "calc".to_owned(),
594 instance_id: "calc-1".to_owned(),
595 endpoint_uri: "http://calc:8080".to_owned(),
596 version: "1.2.3".to_owned(),
597 rest_endpoint_uri: Some("http://calc:8080".to_owned()),
598 openapi_spec_hash: Some("1a2b3c4d5e6f7a8b".to_owned()),
599 labels: [("shard".to_owned(), "7".to_owned())].into_iter().collect(),
600 state: ProtoInstanceState::Healthy as i32,
601 };
602 let domain = proto_instance_to_domain(proto);
603 assert_eq!(domain.gear, "calc");
604 assert_eq!(domain.state, InstanceState::Healthy);
605 assert_eq!(domain.instance_id, "calc-1");
606 assert_eq!(
607 domain.endpoint.as_ref().map(|e| e.uri.as_str()),
608 Some("http://calc:8080")
609 );
610 assert_eq!(domain.version.as_deref(), Some("1.2.3"));
611 assert_eq!(
612 domain.rest_endpoint.map(|e| e.uri),
613 Some("http://calc:8080".to_owned())
614 );
615 assert_eq!(
617 domain.openapi_spec_hash.as_deref(),
618 Some("1a2b3c4d5e6f7a8b")
619 );
620 assert_eq!(domain.labels.get("shard"), Some(&"7".to_owned()));
622 }
623
624 #[test]
625 fn proto_instance_maps_empty_version_to_none() {
626 let proto = InstanceInfo {
627 gear_name: "worker".to_owned(),
628 instance_id: "worker-1".to_owned(),
629 endpoint_uri: "http://worker:7000".to_owned(),
630 version: String::new(),
631 rest_endpoint_uri: None,
632 openapi_spec_hash: None,
633 labels: std::collections::HashMap::new(),
634 state: ProtoInstanceState::Unspecified as i32,
635 };
636 let domain = proto_instance_to_domain(proto);
637 assert_eq!(
639 domain.endpoint.as_ref().map(|e| e.uri.as_str()),
640 Some("http://worker:7000")
641 );
642 assert!(domain.version.is_none());
644 assert!(domain.rest_endpoint.is_none());
645 assert!(domain.openapi_spec_hash.is_none());
646 assert!(domain.labels.is_empty());
647 assert_eq!(domain.state, InstanceState::Unknown);
650 assert!(!domain.state.is_serving());
651 }
652
653 #[test]
654 fn proto_instance_maps_empty_endpoint_to_none() {
655 let proto = InstanceInfo {
659 gear_name: "grpc-only".to_owned(),
660 instance_id: "g-1".to_owned(),
661 endpoint_uri: String::new(),
662 version: String::new(),
663 rest_endpoint_uri: None,
664 openapi_spec_hash: None,
665 labels: std::collections::HashMap::new(),
666 state: ProtoInstanceState::Ready as i32,
667 };
668 let domain = proto_instance_to_domain(proto);
669 assert!(
670 domain.endpoint.is_none(),
671 "an empty proto endpoint_uri must map to None"
672 );
673 }
674
675 #[test]
676 fn proto_state_unspecified_and_unrecognised_map_to_unknown() {
677 assert_eq!(
681 proto_state_to_domain(ProtoInstanceState::Unspecified as i32),
682 InstanceState::Unknown
683 );
684 assert_eq!(proto_state_to_domain(9999), InstanceState::Unknown);
685 assert!(!proto_state_to_domain(9999).is_serving());
686
687 assert_eq!(
689 proto_state_to_domain(ProtoInstanceState::Registered as i32),
690 InstanceState::Registered
691 );
692 assert_eq!(
693 proto_state_to_domain(ProtoInstanceState::Healthy as i32),
694 InstanceState::Healthy
695 );
696 }
697
698 #[test]
699 fn call_error_types_invalid_argument_as_permanent() {
700 let err = call_error(
701 "register_instance",
702 &tonic::Status::invalid_argument("bad label"),
703 );
704 assert!(
705 err.downcast_ref::<DirectoryInvalidArgument>().is_some(),
706 "InvalidArgument must map to the typed permanent DirectoryInvalidArgument"
707 );
708 }
709
710 #[test]
711 fn call_error_types_permission_denied_as_permanent() {
712 let err = call_error(
716 "register_instance",
717 &tonic::Status::permission_denied("peer not authorized"),
718 );
719 assert!(
720 err.downcast_ref::<DirectoryPermissionDenied>().is_some(),
721 "PermissionDenied must map to the typed permanent DirectoryPermissionDenied"
722 );
723 }
724
725 #[test]
726 fn call_error_keeps_transient_codes_opaque() {
727 for status in [
734 tonic::Status::unavailable("connection reset"),
735 tonic::Status::unauthenticated("invalid internal token"),
736 tonic::Status::failed_precondition("service name already owned by another gear"),
737 ] {
738 let err = call_error("heartbeat", &status);
739 assert!(err.downcast_ref::<DirectoryInvalidArgument>().is_none());
740 assert!(
741 err.downcast_ref::<DirectoryPermissionDenied>().is_none(),
742 "{:?} must stay transient/retryable, not map to DirectoryPermissionDenied",
743 status.code()
744 );
745 assert!(err.to_string().contains(&format!("{:?}", status.code())));
746 }
747 }
748}