1use std::sync::Arc;
16use std::sync::atomic::{AtomicBool, Ordering};
17use std::time::Duration;
18
19use api::v1::HealthCheckRequest;
20use api::v1::flow::flow_client::FlowClient as PbFlowClient;
21use api::v1::health_check_client::HealthCheckClient;
22use api::v1::prometheus_gateway_client::PrometheusGatewayClient;
23use api::v1::region::region_client::RegionClient as PbRegionClient;
24use arrow_flight::flight_service_client::FlightServiceClient;
25use common_grpc::channel_manager::{
26 ChannelConfig, ChannelManager, ClientTlsOption, load_client_tls_config,
27};
28use parking_lot::RwLock;
29use snafu::{OptionExt, ResultExt};
30use tonic::codec::CompressionEncoding;
31use tonic::transport::Channel;
32
33use crate::load_balance::{LoadBalance, Loadbalancer};
34use crate::{Result, error};
35
36const DEFAULT_HEALTH_CHECK_INTERVAL: Duration = Duration::from_secs(30);
37const DEFAULT_HEALTH_CHECK_TIMEOUT: Duration = Duration::from_secs(1);
38
39#[derive(Clone, Debug)]
41pub struct ClientOptions {
42 pub health_check_interval: Duration,
44 pub health_check_timeout: Duration,
46}
47
48impl Default for ClientOptions {
49 fn default() -> Self {
50 Self {
51 health_check_interval: DEFAULT_HEALTH_CHECK_INTERVAL,
52 health_check_timeout: DEFAULT_HEALTH_CHECK_TIMEOUT,
53 }
54 }
55}
56
57pub struct FlightClient {
58 addr: String,
59 client: FlightServiceClient<Channel>,
60}
61
62impl FlightClient {
63 pub fn addr(&self) -> &str {
64 &self.addr
65 }
66
67 pub fn mut_inner(&mut self) -> &mut FlightServiceClient<Channel> {
68 &mut self.client
69 }
70}
71
72#[derive(Clone, Debug, Default)]
73pub struct Client {
74 inner: Arc<Inner>,
75}
76
77#[derive(Debug)]
78struct Inner {
79 query_channel_manager: ChannelManager,
80 control_channel_manager: ChannelManager,
81 peers: RwLock<Peers>,
82 load_balance: Loadbalancer,
83 health_check_interval: Duration,
84 health_check_timeout: Duration,
85 health_check_started: AtomicBool,
86}
87
88impl Default for Inner {
89 fn default() -> Self {
90 #[allow(deprecated)]
91 Self::with_manager_and_peers(ChannelManager::new(), Vec::new(), ClientOptions::default())
92 }
93}
94
95#[derive(Debug, Default, PartialEq, Eq)]
96struct PeerStates {
97 active: Vec<usize>,
98 inactive: Vec<usize>,
99}
100
101#[derive(Debug, Default)]
102struct Peers {
103 addresses: Vec<String>,
104 states: PeerStates,
105 generation: u64,
106}
107
108impl Inner {
109 #[deprecated(note = "legacy single-manager path shares its pool between lanes")]
110 fn with_manager_and_peers(
111 channel_manager: ChannelManager,
112 peers: Vec<String>,
113 options: ClientOptions,
114 ) -> Self {
115 Self::with_managers_and_peers(channel_manager.clone(), channel_manager, peers, options)
116 }
117
118 fn with_managers_and_peers(
120 query_channel_manager: ChannelManager,
121 control_channel_manager: ChannelManager,
122 peers: Vec<String>,
123 options: ClientOptions,
124 ) -> Self {
125 let peer_count = peers.len();
126 Self {
127 query_channel_manager,
128 control_channel_manager,
129 peers: RwLock::new(Peers {
130 addresses: peers,
131 states: PeerStates {
132 active: (0..peer_count).collect(),
133 inactive: Vec::new(),
134 },
135 generation: 0,
136 }),
137 load_balance: Loadbalancer::default(),
138 health_check_interval: options.health_check_interval,
139 health_check_timeout: options.health_check_timeout,
140 health_check_started: AtomicBool::new(false),
141 }
142 }
143
144 fn set_peers(&self, addresses: Vec<String>) {
145 let peer_count = addresses.len();
146 let mut peers = self.peers.write();
147 peers.addresses = addresses;
148 peers.states = PeerStates {
149 active: (0..peer_count).collect(),
150 inactive: Vec::new(),
151 };
152 peers.generation = peers.generation.wrapping_add(1);
153 }
154
155 fn peer_count(&self) -> usize {
156 self.peers.read().addresses.len()
157 }
158
159 fn get_peer(&self) -> Option<String> {
160 let peers = self.peers.read();
161 let index = self
162 .load_balance
163 .get_index(&peers.states.active)
164 .or_else(|| self.load_balance.get_index(&peers.states.inactive))?;
165 Some(peers.addresses[*index].clone())
166 }
167
168 async fn refresh_peer_states(&self) {
169 let (generation, peers) = {
170 let peers = self.peers.read();
171 let addresses = peers
172 .states
173 .active
174 .iter()
175 .chain(&peers.states.inactive)
176 .map(|&index| (index, peers.addresses[index].clone()))
177 .collect::<Vec<_>>();
178 (peers.generation, addresses)
179 };
180 let health_checks = peers.into_iter().map(|(index, addr)| async move {
181 let is_active = self.check_peer_health(&addr).await;
182 (index, is_active)
183 });
184 let results = futures::future::join_all(health_checks).await;
185
186 let (active, inactive) = results.into_iter().fold(
187 (Vec::new(), Vec::new()),
188 |(mut active, mut inactive), (index, is_active)| {
189 if is_active {
190 active.push(index);
191 } else {
192 inactive.push(index);
193 }
194 (active, inactive)
195 },
196 );
197
198 let mut peers = self.peers.write();
199 if peers.generation == generation {
200 peers.states = PeerStates { active, inactive };
201 }
202 }
203
204 async fn check_peer_health(&self, addr: &str) -> bool {
205 let Ok(channel) = self.control_channel_manager.get(addr) else {
206 return false;
207 };
208 let mut client = HealthCheckClient::new(channel);
209 tokio::time::timeout(
210 self.health_check_timeout,
211 client.health_check(HealthCheckRequest {}),
212 )
213 .await
214 .is_ok_and(|result| result.is_ok())
215 }
216}
217
218fn random_initial_delay(max_delay: Duration) -> Duration {
219 let max_nanos = max_delay.as_nanos().min(u64::MAX as u128) as u64;
220 if max_nanos == 0 {
221 return Duration::ZERO;
222 }
223
224 Duration::from_nanos(rand::random_range(0..max_nanos))
225}
226
227impl Client {
228 #[deprecated(
230 note = "shares one manager between query and control lanes; use `with_query_and_control_managers` with independently constructed managers instead"
231 )]
232 pub fn new() -> Self {
233 Default::default()
234 }
235
236 #[deprecated(
238 note = "shares one manager between query and control lanes; use `with_query_and_control_managers` with independently constructed managers instead"
239 )]
240 pub fn with_urls<U, A>(urls: A) -> Self
241 where
242 U: AsRef<str>,
243 A: AsRef<[U]>,
244 {
245 #[allow(deprecated)]
246 Self::with_urls_and_options(urls, ClientOptions::default())
247 }
248
249 #[deprecated(
253 note = "shares one manager between query and control lanes; use `with_query_and_control_managers` instead"
254 )]
255 pub fn with_urls_and_options<U, A>(urls: A, options: ClientOptions) -> Self
256 where
257 U: AsRef<str>,
258 A: AsRef<[U]>,
259 {
260 #[allow(deprecated)]
261 Self::with_manager_and_urls_and_options(ChannelManager::new(), urls, options)
262 }
263
264 #[deprecated(
266 note = "shares one manager between query and control lanes; use `with_query_and_control_managers` instead"
267 )]
268 pub fn with_tls_and_urls<U, A>(urls: A, client_tls: ClientTlsOption) -> Result<Self>
269 where
270 U: AsRef<str>,
271 A: AsRef<[U]>,
272 {
273 #[allow(deprecated)]
274 Self::with_tls_and_urls_and_options(urls, client_tls, ClientOptions::default())
275 }
276
277 #[deprecated(
281 note = "shares one manager between query and control lanes; use `with_query_and_control_managers` instead"
282 )]
283 pub fn with_tls_and_urls_and_options<U, A>(
284 urls: A,
285 client_tls: ClientTlsOption,
286 options: ClientOptions,
287 ) -> Result<Self>
288 where
289 U: AsRef<str>,
290 A: AsRef<[U]>,
291 {
292 let channel_config = ChannelConfig::default().client_tls_config(client_tls.clone());
293 let tls_config =
294 load_client_tls_config(Some(client_tls)).context(error::CreateTlsChannelSnafu)?;
295 let channel_manager = ChannelManager::with_config(channel_config, tls_config);
296 #[allow(deprecated)]
297 Ok(Self::with_manager_and_urls_and_options(
298 channel_manager,
299 urls,
300 options,
301 ))
302 }
303
304 #[deprecated(
306 note = "shares one manager between query and control lanes; use `with_query_and_control_managers` with independently constructed managers instead"
307 )]
308 pub fn with_manager_and_urls<U, A>(channel_manager: ChannelManager, urls: A) -> Self
309 where
310 U: AsRef<str>,
311 A: AsRef<[U]>,
312 {
313 #[allow(deprecated)]
314 Self::with_manager_and_urls_and_options(channel_manager, urls, ClientOptions::default())
315 }
316
317 pub fn with_query_and_control_managers<U, A>(
322 query_channel_manager: ChannelManager,
323 control_channel_manager: ChannelManager,
324 urls: A,
325 ) -> Self
326 where
327 U: AsRef<str>,
328 A: AsRef<[U]>,
329 {
330 Self::with_query_and_control_managers_and_options(
331 query_channel_manager,
332 control_channel_manager,
333 urls,
334 ClientOptions::default(),
335 )
336 }
337
338 fn with_query_and_control_managers_and_options<U, A>(
339 query_channel_manager: ChannelManager,
340 control_channel_manager: ChannelManager,
341 urls: A,
342 options: ClientOptions,
343 ) -> Self
344 where
345 U: AsRef<str>,
346 A: AsRef<[U]>,
347 {
348 let urls: Vec<String> = urls
349 .as_ref()
350 .iter()
351 .map(|peer| peer.as_ref().to_string())
352 .collect();
353 Self {
354 inner: Arc::new(Inner::with_managers_and_peers(
355 query_channel_manager,
356 control_channel_manager,
357 urls,
358 options,
359 )),
360 }
361 }
362
363 #[deprecated(
367 note = "shares this manager and its pool between query and control lanes; use `with_query_and_control_managers` instead"
368 )]
369 pub fn with_manager_and_urls_and_options<U, A>(
370 channel_manager: ChannelManager,
371 urls: A,
372 options: ClientOptions,
373 ) -> Self
374 where
375 U: AsRef<str>,
376 A: AsRef<[U]>,
377 {
378 let channel_manager_for_query = channel_manager.clone();
379 Self::with_query_and_control_managers_and_options(
381 channel_manager_for_query,
382 channel_manager,
383 urls,
384 options,
385 )
386 }
387
388 pub fn start<U, A>(&self, urls: A)
389 where
390 U: AsRef<str>,
391 A: AsRef<[U]>,
392 {
393 let urls = urls
394 .as_ref()
395 .iter()
396 .map(|peer| peer.as_ref().to_string())
397 .collect();
398 self.inner.set_peers(urls);
399 }
400
401 fn trigger_health_check(&self) {
402 if self.inner.health_check_interval.is_zero() || self.inner.peer_count() <= 1 {
403 return;
404 }
405
406 if self
407 .inner
408 .health_check_started
409 .swap(true, Ordering::Relaxed)
410 {
411 return;
412 }
413
414 let inner = Arc::downgrade(&self.inner);
415 let health_check_interval = self.inner.health_check_interval;
416 common_runtime::spawn_global(async move {
417 tokio::time::sleep(random_initial_delay(health_check_interval)).await;
418 let mut interval = tokio::time::interval(health_check_interval);
419 interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
420
421 loop {
422 interval.tick().await;
423 let Some(inner) = inner.upgrade() else {
424 return;
425 };
426 if inner.peer_count() > 1 {
427 inner.refresh_peer_states().await;
428 }
429 }
430 });
431 }
432
433 pub fn find_channel(&self) -> Result<(String, Channel)> {
434 self.trigger_health_check();
435
436 let addr = self
437 .inner
438 .get_peer()
439 .context(error::IllegalGrpcClientStateSnafu {
440 err_msg: "No available peer found",
441 })?;
442
443 let channel = self
444 .inner
445 .control_channel_manager
446 .get(&addr)
447 .context(error::CreateChannelSnafu { addr: &addr })?;
448 Ok((addr, channel))
449 }
450
451 pub fn max_grpc_recv_message_size(&self) -> usize {
452 self.inner
453 .control_channel_manager
454 .config()
455 .max_recv_message_size
456 .as_bytes() as usize
457 }
458
459 pub fn max_grpc_send_message_size(&self) -> usize {
460 self.inner
461 .control_channel_manager
462 .config()
463 .max_send_message_size
464 .as_bytes() as usize
465 }
466
467 pub fn make_flight_client(
471 &self,
472 send_compression: bool,
473 accept_compression: bool,
474 ) -> Result<FlightClient> {
475 self.make_flight_client_with_manager(
476 &self.inner.query_channel_manager,
477 send_compression,
478 accept_compression,
479 )
480 }
481
482 pub(crate) fn make_control_flight_client(
483 &self,
484 send_compression: bool,
485 accept_compression: bool,
486 ) -> Result<FlightClient> {
487 self.make_flight_client_with_manager(
488 &self.inner.control_channel_manager,
489 send_compression,
490 accept_compression,
491 )
492 }
493
494 fn make_flight_client_with_manager(
495 &self,
496 channel_manager: &ChannelManager,
497 send_compression: bool,
498 accept_compression: bool,
499 ) -> Result<FlightClient> {
500 self.trigger_health_check();
501 let addr = self
502 .inner
503 .get_peer()
504 .context(error::IllegalGrpcClientStateSnafu {
505 err_msg: "No available peer found",
506 })?;
507 let channel = channel_manager
508 .get(&addr)
509 .context(error::CreateChannelSnafu { addr: &addr })?;
510
511 let mut client = FlightServiceClient::new(channel)
512 .max_decoding_message_size(
513 channel_manager.config().max_recv_message_size.as_bytes() as usize
514 )
515 .max_encoding_message_size(
516 channel_manager.config().max_send_message_size.as_bytes() as usize
517 );
518 if send_compression {
520 client = client.send_compressed(CompressionEncoding::Zstd);
521 }
522 if accept_compression {
523 client = client.accept_compressed(CompressionEncoding::Zstd);
524 }
525
526 Ok(FlightClient { addr, client })
527 }
528
529 pub(crate) fn raw_region_client(&self) -> Result<(String, PbRegionClient<Channel>)> {
530 let (addr, channel) = self.find_channel()?;
531 let client = PbRegionClient::new(channel)
532 .max_decoding_message_size(
533 self.inner
534 .control_channel_manager
535 .config()
536 .max_recv_message_size
537 .as_bytes() as usize,
538 )
539 .max_encoding_message_size(
540 self.inner
541 .control_channel_manager
542 .config()
543 .max_send_message_size
544 .as_bytes() as usize,
545 );
546 Ok((addr, client))
547 }
548
549 pub(crate) fn raw_flow_client(&self) -> Result<(String, PbFlowClient<Channel>)> {
550 let (addr, channel) = self.find_channel()?;
551 let client = PbFlowClient::new(channel)
552 .max_decoding_message_size(
553 self.inner
554 .control_channel_manager
555 .config()
556 .max_recv_message_size
557 .as_bytes() as usize,
558 )
559 .max_encoding_message_size(
560 self.inner
561 .control_channel_manager
562 .config()
563 .max_send_message_size
564 .as_bytes() as usize,
565 )
566 .accept_compressed(CompressionEncoding::Zstd)
567 .send_compressed(CompressionEncoding::Zstd);
568 Ok((addr, client))
569 }
570
571 pub fn make_prometheus_gateway_client(&self) -> Result<PrometheusGatewayClient<Channel>> {
572 let (_, channel) = self.find_channel()?;
573 let client = PrometheusGatewayClient::new(channel)
574 .accept_compressed(CompressionEncoding::Gzip)
575 .accept_compressed(CompressionEncoding::Zstd)
576 .send_compressed(CompressionEncoding::Gzip)
577 .send_compressed(CompressionEncoding::Zstd);
578 Ok(client)
579 }
580
581 pub async fn health_check(&self) -> Result<()> {
582 let (_, channel) = self.find_channel()?;
583 let mut client = HealthCheckClient::new(channel);
584 let _ = client.health_check(HealthCheckRequest {}).await?;
585 Ok(())
586 }
587
588 #[cfg(feature = "testing")]
590 pub fn channel_pool_sizes(&self) -> (usize, usize) {
591 let pool_size = |channel_manager: &ChannelManager| {
592 let mut size = 0;
593 channel_manager.retain_channel(|_, _| {
594 size += 1;
595 true
596 });
597 size
598 };
599 (
600 pool_size(&self.inner.query_channel_manager),
601 pool_size(&self.inner.control_channel_manager),
602 )
603 }
604
605 #[cfg(feature = "testing")]
607 pub fn peer_addresses_by_state(&self) -> (Vec<String>, Vec<String>) {
608 let peers = self.inner.peers.read();
609 let addresses = |indices: &[usize]| {
610 indices
611 .iter()
612 .map(|&index| peers.addresses[index].clone())
613 .collect()
614 };
615 (
616 addresses(&peers.states.active),
617 addresses(&peers.states.inactive),
618 )
619 }
620}
621
622#[cfg(test)]
623#[allow(deprecated)]
624mod tests {
625 use std::collections::HashSet;
626 use std::sync::Arc;
627 use std::sync::atomic::Ordering;
628 use std::time::Duration;
629
630 use api::v1::health_check_server::{HealthCheck, HealthCheckServer};
631 use api::v1::{HealthCheckRequest, HealthCheckResponse};
632 use common_grpc::channel_manager::ChannelManager;
633 use tokio::net::TcpListener;
634 use tokio::sync::Notify;
635 use tokio::task::JoinHandle;
636 use tokio::time::{interval, timeout};
637 use tokio_stream::wrappers::TcpListenerStream;
638 use tonic::{Request, Response, Status};
639
640 use super::{Client, ClientOptions, Inner, PeerStates};
641 use crate::load_balance::Loadbalancer;
642
643 const HEALTH_REFRESH_INTERVAL: Duration = Duration::from_millis(10);
644 const STATE_REFRESH_TIMEOUT: Duration = Duration::from_secs(1);
645
646 struct HealthyHealthCheck;
647
648 #[tonic::async_trait]
649 impl HealthCheck for HealthyHealthCheck {
650 async fn health_check(
651 &self,
652 _request: Request<HealthCheckRequest>,
653 ) -> Result<Response<HealthCheckResponse>, Status> {
654 Ok(Response::new(HealthCheckResponse {}))
655 }
656 }
657
658 struct UnhealthyHealthCheck;
659
660 #[tonic::async_trait]
661 impl HealthCheck for UnhealthyHealthCheck {
662 async fn health_check(
663 &self,
664 _request: Request<HealthCheckRequest>,
665 ) -> Result<Response<HealthCheckResponse>, Status> {
666 Err(Status::unavailable("peer is unavailable"))
667 }
668 }
669
670 struct PendingHealthCheck {
671 started: Option<Arc<Notify>>,
672 }
673
674 #[tonic::async_trait]
675 impl HealthCheck for PendingHealthCheck {
676 async fn health_check(
677 &self,
678 _request: Request<HealthCheckRequest>,
679 ) -> Result<Response<HealthCheckResponse>, Status> {
680 if let Some(started) = &self.started {
681 started.notify_one();
682 }
683 std::future::pending().await
684 }
685 }
686
687 async fn start_health_check_server<T>(handler: T) -> (String, JoinHandle<()>)
688 where
689 T: HealthCheck + Send + Sync + 'static,
690 {
691 let listener = TcpListener::bind("127.0.0.1:0")
692 .await
693 .expect("bind health check server");
694 let addr = listener
695 .local_addr()
696 .expect("read health check server address")
697 .to_string();
698 let server = tokio::spawn(async move {
699 tonic::transport::Server::builder()
700 .add_service(HealthCheckServer::new(handler))
701 .serve_with_incoming(TcpListenerStream::new(listener))
702 .await
703 .expect("serve health check server");
704 });
705
706 (addr, server)
707 }
708
709 async fn wait_for_peer_states(client: &Client, expected: PeerStates) {
710 let mut poll = interval(HEALTH_REFRESH_INTERVAL);
711 timeout(STATE_REFRESH_TIMEOUT, async {
712 loop {
713 poll.tick().await;
714 if client.inner.peers.read().states == expected {
715 return;
716 }
717 }
718 })
719 .await
720 .expect("health refresh did not reach expected peer states");
721 }
722
723 fn mock_peers() -> Vec<String> {
724 vec![
725 "127.0.0.1:3001".to_string(),
726 "127.0.0.1:3002".to_string(),
727 "127.0.0.1:3003".to_string(),
728 ]
729 }
730
731 #[tokio::test]
732 async fn test_explicit_dual_manager_constructor_uses_isolated_reused_channel_pools() {
733 let query_channel_manager = ChannelManager::new();
734 let control_channel_manager = ChannelManager::new();
735 let client = Client::with_query_and_control_managers(
736 query_channel_manager.clone(),
737 control_channel_manager.clone(),
738 ["127.0.0.1:3001"],
739 );
740
741 client.make_flight_client(false, false).unwrap();
742 client.make_flight_client(false, false).unwrap();
743 assert_eq!(1, channel_pool_size(&query_channel_manager));
744 assert_eq!(0, channel_pool_size(&control_channel_manager));
745
746 client.make_control_flight_client(false, false).unwrap();
747 client.make_control_flight_client(false, false).unwrap();
748 assert_eq!(1, channel_pool_size(&query_channel_manager));
749 assert_eq!(1, channel_pool_size(&control_channel_manager));
750 }
751
752 fn channel_pool_size(manager: &ChannelManager) -> usize {
753 let mut size = 0;
754 manager.retain_channel(|_, _| {
755 size += 1;
756 true
757 });
758 size
759 }
760
761 #[test]
762 fn test_inner() {
763 let inner = Inner::default();
764
765 assert!(matches!(
766 inner.load_balance,
767 Loadbalancer::Random(crate::load_balance::Random)
768 ));
769 assert!(inner.get_peer().is_none());
770
771 let peers = mock_peers();
772 let all: HashSet<String> = peers.iter().cloned().collect();
773 let inner =
774 Inner::with_manager_and_peers(ChannelManager::new(), peers, ClientOptions::default());
775
776 for _ in 0..20 {
777 assert!(all.contains(&inner.get_peer().unwrap()));
778 }
779 }
780
781 #[test]
782 fn test_inner_prefers_active_peer() {
783 let peers = mock_peers();
784 let inner = Inner::with_manager_and_peers(
785 ChannelManager::new(),
786 peers.clone(),
787 ClientOptions::default(),
788 );
789 inner.peers.write().states = PeerStates {
790 active: vec![0],
791 inactive: vec![1, 2],
792 };
793
794 assert_eq!(Some(peers[0].clone()), inner.get_peer());
795 }
796
797 #[test]
798 fn test_zero_health_check_interval_disables_background_task() {
799 let client = Client::with_urls_and_options(
800 mock_peers(),
801 ClientOptions {
802 health_check_interval: Duration::ZERO,
803 ..Default::default()
804 },
805 );
806
807 assert!(!client.inner.health_check_started.load(Ordering::Relaxed));
808 let peers = client.inner.peers.read();
809 assert_eq!(mock_peers(), peers.addresses);
810 assert_eq!(vec![0, 1, 2], peers.states.active);
811 assert!(peers.states.inactive.is_empty());
812 }
813
814 #[test]
815 fn test_multi_peer_constructor_defers_background_task() {
816 let client = Client::with_urls(mock_peers());
817
818 assert!(!client.inner.health_check_started.load(Ordering::Relaxed));
819 }
820
821 #[tokio::test]
822 async fn test_single_peer_does_not_start_background_task() {
823 let client = Client::with_urls(["127.0.0.1:3001"]);
824
825 client.find_channel().unwrap();
826
827 assert!(!client.inner.health_check_started.load(Ordering::Relaxed));
828 }
829
830 #[test]
831 fn test_start_initializes_new_client_without_starting_background_task() {
832 let client = Client::new();
833 let peers = mock_peers();
834
835 client.start(peers.clone());
836
837 assert!(peers.contains(&client.inner.get_peer().unwrap()));
838 assert!(!client.inner.health_check_started.load(Ordering::Relaxed));
839 }
840
841 #[tokio::test]
842 async fn test_health_refresh_marks_unhealthy_peer_inactive_and_selects_healthy_peer() {
843 let (healthy_addr, healthy_server) = start_health_check_server(HealthyHealthCheck).await;
845 let (unhealthy_addr, unhealthy_server) =
846 start_health_check_server(UnhealthyHealthCheck).await;
847 let client = Client::with_urls_and_options(
848 [healthy_addr.clone(), unhealthy_addr],
849 ClientOptions {
850 health_check_interval: HEALTH_REFRESH_INTERVAL,
851 ..Default::default()
852 },
853 );
854 assert!(!client.inner.health_check_started.load(Ordering::Relaxed));
855
856 client.find_channel().unwrap();
858 assert!(client.inner.health_check_started.load(Ordering::Relaxed));
859 wait_for_peer_states(
860 &client,
861 PeerStates {
862 active: vec![0],
863 inactive: vec![1],
864 },
865 )
866 .await;
867
868 assert_eq!(Some(healthy_addr), client.inner.get_peer());
870
871 healthy_server.abort();
872 unhealthy_server.abort();
873 }
874
875 #[tokio::test]
876 async fn test_health_refresh_times_out_pending_peer() {
877 let (healthy_addr, healthy_server) = start_health_check_server(HealthyHealthCheck).await;
878 let (pending_addr, pending_server) =
879 start_health_check_server(PendingHealthCheck { started: None }).await;
880 let client = Client::with_urls_and_options(
881 [healthy_addr, pending_addr],
882 ClientOptions {
883 health_check_interval: HEALTH_REFRESH_INTERVAL,
884 health_check_timeout: Duration::from_millis(20),
885 },
886 );
887
888 client.find_channel().unwrap();
889 wait_for_peer_states(
890 &client,
891 PeerStates {
892 active: vec![0],
893 inactive: vec![1],
894 },
895 )
896 .await;
897
898 healthy_server.abort();
899 pending_server.abort();
900 }
901
902 #[tokio::test]
903 async fn test_peer_update_ignores_in_flight_health_result() {
904 let started = Arc::new(Notify::new());
905 let (pending_addr, pending_server) = start_health_check_server(PendingHealthCheck {
906 started: Some(started.clone()),
907 })
908 .await;
909 let inner = Arc::new(Inner::with_manager_and_peers(
910 ChannelManager::new(),
911 vec![pending_addr],
912 ClientOptions {
913 health_check_timeout: Duration::from_millis(20),
914 ..Default::default()
915 },
916 ));
917 let refresh_inner = inner.clone();
918 let refresh = tokio::spawn(async move {
919 refresh_inner.refresh_peer_states().await;
920 });
921 timeout(STATE_REFRESH_TIMEOUT, started.notified())
922 .await
923 .expect("pending health check did not start");
924
925 inner.set_peers(vec!["127.0.0.1:3001".to_string()]);
926 refresh.await.unwrap();
927
928 let peers = inner.peers.read();
929 assert_eq!(vec!["127.0.0.1:3001"], peers.addresses);
930 assert_eq!(vec![0], peers.states.active);
931 assert!(peers.states.inactive.is_empty());
932 pending_server.abort();
933 }
934}