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, InstanceState, LabelSelector,
13 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 {
60 match status.code() {
61 tonic::Code::InvalidArgument => {
62 DirectoryInvalidArgument::new(status.message().to_owned()).into()
63 }
64 code => anyhow::anyhow!("directory {op} failed: gRPC {code:?}: {}", status.message()),
65 }
66}
67
68pub struct DirectoryGrpcClient {
79 inner: DirectoryServiceClient<AuthedChannel>,
80}
81
82impl DirectoryGrpcClient {
83 pub async fn connect(uri: impl Into<String>) -> Result<Self> {
91 let cfg = GrpcClientConfig::new("directory");
92 Self::connect_with_retry(uri, &cfg).await
93 }
94
95 pub async fn connect_with_interceptor(
107 uri: impl Into<String>,
108 interceptor: InternalAuthInterceptor,
109 ) -> Result<Self> {
110 let cfg = GrpcClientConfig::new("directory");
111 let channel: Channel = connect_with_retry(uri, &cfg).await?;
112 Ok(Self::from_channel_with_interceptor(channel, interceptor))
113 }
114
115 pub fn connect_lazy(uri: impl Into<String>) -> Result<Self> {
134 let cfg = GrpcClientConfig::new("directory");
135 let channel: Channel = connect_lazy(uri, &cfg)?;
136 Ok(Self::from_channel(channel))
137 }
138
139 pub fn connect_lazy_with_interceptor(
151 uri: impl Into<String>,
152 interceptor: InternalAuthInterceptor,
153 ) -> Result<Self> {
154 let cfg = GrpcClientConfig::new("directory");
155 let channel: Channel = connect_lazy(uri, &cfg)?;
157 Ok(Self::from_channel_with_interceptor(channel, interceptor))
158 }
159
160 pub async fn connect_with_retry(
168 uri: impl Into<String>,
169 cfg: &GrpcClientConfig,
170 ) -> Result<Self> {
171 let channel: Channel = connect_with_retry(uri, cfg).await?;
172 Ok(Self::from_channel(channel))
173 }
174
175 pub async fn connect_no_retry(uri: impl Into<String>, cfg: &GrpcClientConfig) -> Result<Self> {
183 let uri_string = uri.into();
184
185 let endpoint = tonic::transport::Endpoint::from_shared(uri_string)?
187 .connect_timeout(cfg.connect_timeout)
188 .timeout(cfg.rpc_timeout);
189
190 let channel = endpoint.connect().await?;
192
193 if cfg.enable_tracing {
194 tracing::debug!(
195 service_name = cfg.service_name,
196 connect_timeout_ms = cfg.connect_timeout.as_millis(),
197 rpc_timeout_ms = cfg.rpc_timeout.as_millis(),
198 "directory gRPC client connected"
199 );
200 }
201
202 Ok(Self::from_channel(channel))
203 }
204
205 #[must_use]
211 pub fn from_channel(channel: Channel) -> Self {
212 Self::from_channel_with_interceptor(channel, InternalAuthInterceptor::disabled())
213 }
214
215 #[must_use]
218 pub fn from_channel_with_interceptor(
219 channel: Channel,
220 interceptor: InternalAuthInterceptor,
221 ) -> Self {
222 Self {
223 inner: DirectoryServiceClient::with_interceptor(channel, interceptor),
224 }
225 }
226}
227
228#[async_trait]
229impl DirectoryClient for DirectoryGrpcClient {
230 async fn resolve_grpc_service(&self, service_name: &str) -> Result<ServiceEndpoint> {
231 let mut client = self.inner.clone();
232 let request = tonic::Request::new(ResolveGrpcServiceRequest {
233 service_name: service_name.to_owned(),
234 });
235
236 let response = client
237 .resolve_grpc_service(request)
238 .await
239 .map_err(|e| lookup_error(&format!("service {service_name}"), &e))?;
240
241 let proto_response = response.into_inner();
242 Ok(ServiceEndpoint::new(proto_response.endpoint_uri))
243 }
244
245 async fn resolve_rest_service(&self, gear_name: &str) -> Result<ServiceEndpoint> {
246 let mut client = self.inner.clone();
247 let request = tonic::Request::new(ResolveRestServiceRequest {
248 gear_name: gear_name.to_owned(),
249 });
250
251 let response = client
252 .resolve_rest_service(request)
253 .await
254 .map_err(|e| lookup_error(&format!("gear {gear_name}"), &e))?;
255
256 let proto_response = response.into_inner();
257 Ok(ServiceEndpoint::new(proto_response.endpoint_uri))
258 }
259
260 async fn get_openapi_spec(&self, gear_name: &str) -> Result<String> {
261 let mut client = self.inner.clone();
262 let request = tonic::Request::new(GetOpenApiSpecRequest {
263 gear_name: gear_name.to_owned(),
264 });
265
266 let response = client
267 .get_open_api_spec(request)
268 .await
269 .map_err(|e| lookup_error(&format!("openapi spec for gear {gear_name}"), &e))?;
270
271 Ok(response.into_inner().openapi_spec)
272 }
273
274 async fn list_instances(&self, gear: &str) -> Result<Vec<ServiceInstanceInfo>> {
275 let mut client = self.inner.clone();
276 let request = tonic::Request::new(ListInstancesRequest {
277 gear_name: gear.to_owned(),
278 match_labels: std::collections::HashMap::new(),
279 });
280
281 let response = client
282 .list_instances(request)
283 .await
284 .map_err(|e| lookup_error(&format!("instances of gear {gear}"), &e))?;
285
286 let instances = response
287 .into_inner()
288 .instances
289 .into_iter()
290 .map(proto_instance_to_domain)
291 .collect();
292
293 Ok(instances)
294 }
295
296 async fn resolve_by_labels(
297 &self,
298 gear: &str,
299 selector: &LabelSelector,
300 ) -> Result<Vec<ServiceInstanceInfo>> {
301 let mut client = self.inner.clone();
311 let match_labels = selector
312 .match_labels
313 .iter()
314 .map(|(k, v)| (k.clone(), v.clone()))
315 .collect();
316 let request = tonic::Request::new(ListInstancesRequest {
317 gear_name: gear.to_owned(),
318 match_labels,
319 });
320
321 let response = client
322 .list_instances(request)
323 .await
324 .map_err(|e| lookup_error(&format!("instances of gear {gear}"), &e))?;
325
326 let instances = response
327 .into_inner()
328 .instances
329 .into_iter()
330 .map(proto_instance_to_domain)
331 .filter(|i| selector.matches(&i.labels))
332 .collect();
333
334 Ok(instances)
335 }
336
337 async fn list_all_instances(&self) -> Result<Vec<ServiceInstanceInfo>> {
338 let mut client = self.inner.clone();
339 let response = client
340 .list_all_instances(tonic::Request::new(ListAllInstancesRequest {}))
341 .await
342 .map_err(|e| lookup_error("all instances", &e))?;
343
344 let instances = response
345 .into_inner()
346 .instances
347 .into_iter()
348 .map(|proto| {
349 let mut info = proto_instance_to_domain(proto).without_labels();
350 info.openapi_spec = None;
351 info
352 })
353 .collect();
354
355 Ok(instances)
356 }
357
358 async fn register_instance(&self, info: RegisterInstanceInfo) -> Result<()> {
359 let mut client = self.inner.clone();
360
361 let grpc_services = info
363 .grpc_services
364 .into_iter()
365 .map(|(name, ep)| GrpcServiceEndpoint {
366 service_name: name,
367 endpoint_uri: ep.uri,
368 })
369 .collect();
370
371 let req = RegisterInstanceRequest {
372 gear_name: info.gear,
373 instance_id: info.instance_id,
374 grpc_services,
375 version: info.version.unwrap_or_default(),
376 rest_endpoint_uri: info.rest_endpoint.map(|ep| ep.uri),
377 openapi_spec: info.openapi_spec,
378 labels: info.labels.into_iter().collect(),
379 };
380
381 client
382 .register_instance(tonic::Request::new(req))
383 .await
384 .map_err(|e| call_error("register_instance", &e))?;
385
386 Ok(())
387 }
388
389 async fn deregister_instance(&self, gear: &str, instance_id: &str) -> Result<()> {
390 let mut client = self.inner.clone();
391
392 let req = DeregisterInstanceRequest {
393 gear_name: gear.to_owned(),
394 instance_id: instance_id.to_owned(),
395 };
396
397 client
398 .deregister_instance(tonic::Request::new(req))
399 .await
400 .map_err(|e| call_error("deregister_instance", &e))?;
401
402 Ok(())
403 }
404
405 async fn send_heartbeat(&self, gear: &str, instance_id: &str) -> Result<()> {
406 let mut client = self.inner.clone();
407
408 let req = HeartbeatRequest {
409 gear_name: gear.to_owned(),
410 instance_id: instance_id.to_owned(),
411 };
412
413 client
414 .heartbeat(tonic::Request::new(req))
415 .await
416 .map_err(|e| call_error("heartbeat", &e))?;
417
418 Ok(())
419 }
420}
421
422fn proto_instance_to_domain(proto: InstanceInfo) -> ServiceInstanceInfo {
424 ServiceInstanceInfo {
425 gear: proto.gear_name,
426 instance_id: proto.instance_id,
427 endpoint: if proto.endpoint_uri.is_empty() {
428 None
429 } else {
430 Some(ServiceEndpoint::new(proto.endpoint_uri))
431 },
432 version: if proto.version.is_empty() {
433 None
434 } else {
435 Some(proto.version)
436 },
437 rest_endpoint: proto.rest_endpoint_uri.map(ServiceEndpoint::new),
438 openapi_spec: proto.openapi_spec,
439 openapi_spec_hash: proto.openapi_spec_hash,
440 grpc_services: Vec::new(),
449 labels: proto.labels.into_iter().collect::<BTreeMap<_, _>>(),
452 state: proto_state_to_domain(proto.state),
454 }
455}
456
457fn proto_state_to_domain(state: i32) -> InstanceState {
467 match ProtoInstanceState::try_from(state) {
468 Ok(ProtoInstanceState::Ready) => InstanceState::Ready,
469 Ok(ProtoInstanceState::Healthy) => InstanceState::Healthy,
470 Ok(ProtoInstanceState::Quarantined) => InstanceState::Quarantined,
471 Ok(ProtoInstanceState::Draining) => InstanceState::Draining,
472 Ok(ProtoInstanceState::Registered) => InstanceState::Registered,
473 Ok(ProtoInstanceState::Unspecified) => InstanceState::Unknown,
474 Err(_) => {
475 tracing::warn!(
476 raw_state = state,
477 "directory returned an unrecognised InstanceState discriminant; \
478 treating as Unknown (non-serving)"
479 );
480 InstanceState::Unknown
481 }
482 }
483}
484
485#[cfg(test)]
486#[cfg_attr(coverage_nightly, coverage(off))]
487mod tests {
488 use super::*;
489
490 #[tokio::test]
491 async fn test_grpc_client_can_be_constructed() {
492 let endpoint = tonic::transport::Endpoint::from_static("http://[::1]:50051");
494
495 let channel_result = endpoint.connect().await;
498
499 if let Ok(channel) = channel_result {
501 let _client = DirectoryGrpcClient::from_channel(channel);
502 }
503 }
504
505 #[tokio::test]
506 async fn from_channel_constructs_without_connecting() {
507 let channel = Channel::from_static("http://[::1]:50051").connect_lazy();
511 let _default = DirectoryGrpcClient::from_channel(channel.clone());
512 let _authed = DirectoryGrpcClient::from_channel_with_interceptor(
513 channel,
514 InternalAuthInterceptor::disabled(),
515 );
516 }
517
518 #[tokio::test]
519 async fn connect_lazy_succeeds_against_unreachable_peer() {
520 let plain = DirectoryGrpcClient::connect_lazy("http://127.0.0.1:1");
525 assert!(
526 plain.is_ok(),
527 "connect_lazy must not eagerly connect (unreachable peer -> Ok)"
528 );
529
530 let authed = DirectoryGrpcClient::connect_lazy_with_interceptor(
531 "http://127.0.0.1:1",
532 InternalAuthInterceptor::disabled(),
533 );
534 assert!(
535 authed.is_ok(),
536 "connect_lazy_with_interceptor must not eagerly connect (unreachable peer -> Ok)"
537 );
538 }
539
540 #[tokio::test]
541 async fn connect_lazy_rejects_malformed_uri() {
542 assert!(
545 DirectoryGrpcClient::connect_lazy(String::new()).is_err(),
546 "connect_lazy should fail on a malformed URI"
547 );
548 assert!(
549 DirectoryGrpcClient::connect_lazy_with_interceptor(
550 String::new(),
551 InternalAuthInterceptor::disabled(),
552 )
553 .is_err(),
554 "connect_lazy_with_interceptor should fail on a malformed URI"
555 );
556 }
557
558 #[tokio::test]
559 async fn resolve_grpc_service_through_lazy_client_errors_not_hangs() {
560 let client =
564 DirectoryGrpcClient::connect_lazy("http://127.0.0.1:1").expect("lazy build ok");
565 let outcome = tokio::time::timeout(
566 std::time::Duration::from_secs(5),
567 client.resolve_grpc_service("cf.directory.v1.DirectoryService"),
568 )
569 .await;
570 assert!(
571 outcome.is_ok(),
572 "resolve_grpc_service through a lazy client must not hang against an unreachable peer"
573 );
574 assert!(
575 outcome.unwrap().is_err(),
576 "resolve_grpc_service against an unreachable directory must return Err"
577 );
578 }
579
580 #[test]
581 fn proto_instance_maps_all_fields_to_domain() {
582 let proto = InstanceInfo {
583 gear_name: "calc".to_owned(),
584 instance_id: "calc-1".to_owned(),
585 endpoint_uri: "http://calc:8080".to_owned(),
586 version: "1.2.3".to_owned(),
587 rest_endpoint_uri: Some("http://calc:8080".to_owned()),
588 openapi_spec: Some("{\"openapi\":\"3.1.0\"}".to_owned()),
589 openapi_spec_hash: None,
590 labels: [("shard".to_owned(), "7".to_owned())].into_iter().collect(),
591 state: ProtoInstanceState::Healthy as i32,
592 };
593 let domain = proto_instance_to_domain(proto);
594 assert_eq!(domain.gear, "calc");
595 assert_eq!(domain.state, InstanceState::Healthy);
596 assert_eq!(domain.instance_id, "calc-1");
597 assert_eq!(
598 domain.endpoint.as_ref().map(|e| e.uri.as_str()),
599 Some("http://calc:8080")
600 );
601 assert_eq!(domain.version.as_deref(), Some("1.2.3"));
602 assert_eq!(
603 domain.rest_endpoint.map(|e| e.uri),
604 Some("http://calc:8080".to_owned())
605 );
606 assert!(domain.openapi_spec.is_some());
607 assert_eq!(domain.labels.get("shard"), Some(&"7".to_owned()));
609 }
610
611 #[test]
612 fn proto_instance_maps_empty_version_to_none() {
613 let proto = InstanceInfo {
614 gear_name: "worker".to_owned(),
615 instance_id: "worker-1".to_owned(),
616 endpoint_uri: "http://worker:7000".to_owned(),
617 version: String::new(),
618 rest_endpoint_uri: None,
619 openapi_spec: None,
620 openapi_spec_hash: None,
621 labels: std::collections::HashMap::new(),
622 state: ProtoInstanceState::Unspecified as i32,
623 };
624 let domain = proto_instance_to_domain(proto);
625 assert_eq!(
627 domain.endpoint.as_ref().map(|e| e.uri.as_str()),
628 Some("http://worker:7000")
629 );
630 assert!(domain.version.is_none());
632 assert!(domain.rest_endpoint.is_none());
633 assert!(domain.openapi_spec.is_none());
634 assert!(domain.labels.is_empty());
635 assert_eq!(domain.state, InstanceState::Unknown);
638 assert!(!domain.state.is_serving());
639 }
640
641 #[test]
642 fn proto_instance_maps_empty_endpoint_to_none() {
643 let proto = InstanceInfo {
647 gear_name: "grpc-only".to_owned(),
648 instance_id: "g-1".to_owned(),
649 endpoint_uri: String::new(),
650 version: String::new(),
651 rest_endpoint_uri: None,
652 openapi_spec: None,
653 openapi_spec_hash: None,
654 labels: std::collections::HashMap::new(),
655 state: ProtoInstanceState::Ready as i32,
656 };
657 let domain = proto_instance_to_domain(proto);
658 assert!(
659 domain.endpoint.is_none(),
660 "an empty proto endpoint_uri must map to None"
661 );
662 }
663
664 #[test]
665 fn proto_state_unspecified_and_unrecognised_map_to_unknown() {
666 assert_eq!(
670 proto_state_to_domain(ProtoInstanceState::Unspecified as i32),
671 InstanceState::Unknown
672 );
673 assert_eq!(proto_state_to_domain(9999), InstanceState::Unknown);
674 assert!(!proto_state_to_domain(9999).is_serving());
675
676 assert_eq!(
678 proto_state_to_domain(ProtoInstanceState::Registered as i32),
679 InstanceState::Registered
680 );
681 assert_eq!(
682 proto_state_to_domain(ProtoInstanceState::Healthy as i32),
683 InstanceState::Healthy
684 );
685 }
686}