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| proto_instance_to_domain(proto).without_labels())
349 .collect();
350
351 Ok(instances)
352 }
353
354 async fn register_instance(&self, info: RegisterInstanceInfo) -> Result<()> {
355 let mut client = self.inner.clone();
356
357 let grpc_services = info
359 .grpc_services
360 .into_iter()
361 .map(|(name, ep)| GrpcServiceEndpoint {
362 service_name: name,
363 endpoint_uri: ep.uri,
364 })
365 .collect();
366
367 let req = RegisterInstanceRequest {
368 gear_name: info.gear,
369 instance_id: info.instance_id,
370 grpc_services,
371 version: info.version.unwrap_or_default(),
372 rest_endpoint_uri: info.rest_endpoint.map(|ep| ep.uri),
373 openapi_spec: info.openapi_spec,
374 labels: info.labels.into_iter().collect(),
375 };
376
377 client
378 .register_instance(tonic::Request::new(req))
379 .await
380 .map_err(|e| call_error("register_instance", &e))?;
381
382 Ok(())
383 }
384
385 async fn deregister_instance(&self, gear: &str, instance_id: &str) -> Result<()> {
386 let mut client = self.inner.clone();
387
388 let req = DeregisterInstanceRequest {
389 gear_name: gear.to_owned(),
390 instance_id: instance_id.to_owned(),
391 };
392
393 client
394 .deregister_instance(tonic::Request::new(req))
395 .await
396 .map_err(|e| call_error("deregister_instance", &e))?;
397
398 Ok(())
399 }
400
401 async fn send_heartbeat(&self, gear: &str, instance_id: &str) -> Result<()> {
402 let mut client = self.inner.clone();
403
404 let req = HeartbeatRequest {
405 gear_name: gear.to_owned(),
406 instance_id: instance_id.to_owned(),
407 };
408
409 client
410 .heartbeat(tonic::Request::new(req))
411 .await
412 .map_err(|e| call_error("heartbeat", &e))?;
413
414 Ok(())
415 }
416}
417
418fn proto_instance_to_domain(proto: InstanceInfo) -> ServiceInstanceInfo {
420 ServiceInstanceInfo {
421 gear: proto.gear_name,
422 instance_id: proto.instance_id,
423 endpoint: if proto.endpoint_uri.is_empty() {
424 None
425 } else {
426 Some(ServiceEndpoint::new(proto.endpoint_uri))
427 },
428 version: if proto.version.is_empty() {
429 None
430 } else {
431 Some(proto.version)
432 },
433 rest_endpoint: proto.rest_endpoint_uri.map(ServiceEndpoint::new),
434 openapi_spec_hash: proto.openapi_spec_hash,
435 grpc_services: Vec::new(),
444 labels: proto.labels.into_iter().collect::<BTreeMap<_, _>>(),
447 state: proto_state_to_domain(proto.state),
449 }
450}
451
452fn proto_state_to_domain(state: i32) -> InstanceState {
462 match ProtoInstanceState::try_from(state) {
463 Ok(ProtoInstanceState::Ready) => InstanceState::Ready,
464 Ok(ProtoInstanceState::Healthy) => InstanceState::Healthy,
465 Ok(ProtoInstanceState::Quarantined) => InstanceState::Quarantined,
466 Ok(ProtoInstanceState::Draining) => InstanceState::Draining,
467 Ok(ProtoInstanceState::Registered) => InstanceState::Registered,
468 Ok(ProtoInstanceState::Unspecified) => InstanceState::Unknown,
469 Err(_) => {
470 tracing::warn!(
471 raw_state = state,
472 "directory returned an unrecognised InstanceState discriminant; \
473 treating as Unknown (non-serving)"
474 );
475 InstanceState::Unknown
476 }
477 }
478}
479
480#[cfg(test)]
481#[cfg_attr(coverage_nightly, coverage(off))]
482mod tests {
483 use super::*;
484
485 #[tokio::test]
486 async fn test_grpc_client_can_be_constructed() {
487 let endpoint = tonic::transport::Endpoint::from_static("http://[::1]:50051");
489
490 let channel_result = endpoint.connect().await;
493
494 if let Ok(channel) = channel_result {
496 let _client = DirectoryGrpcClient::from_channel(channel);
497 }
498 }
499
500 #[tokio::test]
501 async fn from_channel_constructs_without_connecting() {
502 let channel = Channel::from_static("http://[::1]:50051").connect_lazy();
506 let _default = DirectoryGrpcClient::from_channel(channel.clone());
507 let _authed = DirectoryGrpcClient::from_channel_with_interceptor(
508 channel,
509 InternalAuthInterceptor::disabled(),
510 );
511 }
512
513 #[tokio::test]
514 async fn connect_lazy_succeeds_against_unreachable_peer() {
515 let plain = DirectoryGrpcClient::connect_lazy("http://127.0.0.1:1");
520 assert!(
521 plain.is_ok(),
522 "connect_lazy must not eagerly connect (unreachable peer -> Ok)"
523 );
524
525 let authed = DirectoryGrpcClient::connect_lazy_with_interceptor(
526 "http://127.0.0.1:1",
527 InternalAuthInterceptor::disabled(),
528 );
529 assert!(
530 authed.is_ok(),
531 "connect_lazy_with_interceptor must not eagerly connect (unreachable peer -> Ok)"
532 );
533 }
534
535 #[tokio::test]
536 async fn connect_lazy_rejects_malformed_uri() {
537 assert!(
540 DirectoryGrpcClient::connect_lazy(String::new()).is_err(),
541 "connect_lazy should fail on a malformed URI"
542 );
543 assert!(
544 DirectoryGrpcClient::connect_lazy_with_interceptor(
545 String::new(),
546 InternalAuthInterceptor::disabled(),
547 )
548 .is_err(),
549 "connect_lazy_with_interceptor should fail on a malformed URI"
550 );
551 }
552
553 #[tokio::test]
554 async fn resolve_grpc_service_through_lazy_client_errors_not_hangs() {
555 let client =
559 DirectoryGrpcClient::connect_lazy("http://127.0.0.1:1").expect("lazy build ok");
560 let outcome = tokio::time::timeout(
561 std::time::Duration::from_secs(5),
562 client.resolve_grpc_service("cf.directory.v1.DirectoryService"),
563 )
564 .await;
565 assert!(
566 outcome.is_ok(),
567 "resolve_grpc_service through a lazy client must not hang against an unreachable peer"
568 );
569 assert!(
570 outcome.unwrap().is_err(),
571 "resolve_grpc_service against an unreachable directory must return Err"
572 );
573 }
574
575 #[test]
576 fn proto_instance_maps_all_fields_to_domain() {
577 let proto = InstanceInfo {
578 gear_name: "calc".to_owned(),
579 instance_id: "calc-1".to_owned(),
580 endpoint_uri: "http://calc:8080".to_owned(),
581 version: "1.2.3".to_owned(),
582 rest_endpoint_uri: Some("http://calc:8080".to_owned()),
583 openapi_spec_hash: Some("1a2b3c4d5e6f7a8b".to_owned()),
584 labels: [("shard".to_owned(), "7".to_owned())].into_iter().collect(),
585 state: ProtoInstanceState::Healthy as i32,
586 };
587 let domain = proto_instance_to_domain(proto);
588 assert_eq!(domain.gear, "calc");
589 assert_eq!(domain.state, InstanceState::Healthy);
590 assert_eq!(domain.instance_id, "calc-1");
591 assert_eq!(
592 domain.endpoint.as_ref().map(|e| e.uri.as_str()),
593 Some("http://calc:8080")
594 );
595 assert_eq!(domain.version.as_deref(), Some("1.2.3"));
596 assert_eq!(
597 domain.rest_endpoint.map(|e| e.uri),
598 Some("http://calc:8080".to_owned())
599 );
600 assert_eq!(
602 domain.openapi_spec_hash.as_deref(),
603 Some("1a2b3c4d5e6f7a8b")
604 );
605 assert_eq!(domain.labels.get("shard"), Some(&"7".to_owned()));
607 }
608
609 #[test]
610 fn proto_instance_maps_empty_version_to_none() {
611 let proto = InstanceInfo {
612 gear_name: "worker".to_owned(),
613 instance_id: "worker-1".to_owned(),
614 endpoint_uri: "http://worker:7000".to_owned(),
615 version: String::new(),
616 rest_endpoint_uri: None,
617 openapi_spec_hash: None,
618 labels: std::collections::HashMap::new(),
619 state: ProtoInstanceState::Unspecified as i32,
620 };
621 let domain = proto_instance_to_domain(proto);
622 assert_eq!(
624 domain.endpoint.as_ref().map(|e| e.uri.as_str()),
625 Some("http://worker:7000")
626 );
627 assert!(domain.version.is_none());
629 assert!(domain.rest_endpoint.is_none());
630 assert!(domain.openapi_spec_hash.is_none());
631 assert!(domain.labels.is_empty());
632 assert_eq!(domain.state, InstanceState::Unknown);
635 assert!(!domain.state.is_serving());
636 }
637
638 #[test]
639 fn proto_instance_maps_empty_endpoint_to_none() {
640 let proto = InstanceInfo {
644 gear_name: "grpc-only".to_owned(),
645 instance_id: "g-1".to_owned(),
646 endpoint_uri: String::new(),
647 version: String::new(),
648 rest_endpoint_uri: None,
649 openapi_spec_hash: None,
650 labels: std::collections::HashMap::new(),
651 state: ProtoInstanceState::Ready as i32,
652 };
653 let domain = proto_instance_to_domain(proto);
654 assert!(
655 domain.endpoint.is_none(),
656 "an empty proto endpoint_uri must map to None"
657 );
658 }
659
660 #[test]
661 fn proto_state_unspecified_and_unrecognised_map_to_unknown() {
662 assert_eq!(
666 proto_state_to_domain(ProtoInstanceState::Unspecified as i32),
667 InstanceState::Unknown
668 );
669 assert_eq!(proto_state_to_domain(9999), InstanceState::Unknown);
670 assert!(!proto_state_to_domain(9999).is_serving());
671
672 assert_eq!(
674 proto_state_to_domain(ProtoInstanceState::Registered as i32),
675 InstanceState::Registered
676 );
677 assert_eq!(
678 proto_state_to_domain(ProtoInstanceState::Healthy as i32),
679 InstanceState::Healthy
680 );
681 }
682}