From e97a6eda030b1922b2cb4b6cb91c2f418686e2e2 Mon Sep 17 00:00:00 2001 From: Michael Ingley Date: Thu, 30 Jul 2026 22:38:59 -0500 Subject: [PATCH] Wire A32 circuit breaking into xDS channels --- tonic-xds/src/client/channel.rs | 275 ++++++- tonic-xds/src/client/circuit_breaking.rs | 767 ++++++++++++++++-- .../src/client/loadbalance/loadbalancer.rs | 46 +- tonic-xds/src/xds/cache.rs | 274 ++++++- tonic-xds/src/xds/cluster_discovery.rs | 75 +- tonic-xds/src/xds/endpoint_manager.rs | 193 ++++- .../src/xds/resource/circuit_breaking.rs | 4 +- tonic-xds/src/xds/resource/cluster.rs | 34 + 8 files changed, 1490 insertions(+), 178 deletions(-) diff --git a/tonic-xds/src/client/channel.rs b/tonic-xds/src/client/channel.rs index 742463a66..44224d88e 100644 --- a/tonic-xds/src/client/channel.rs +++ b/tonic-xds/src/client/channel.rs @@ -22,6 +22,9 @@ * */ +use crate::client::circuit_breaking::{ + CircuitBreakingClusterDiscovery, ClusterCircuitBreakerRegistry, +}; use crate::client::cluster::ClusterClientRegistryGrpc; use crate::client::endpoint::{EndpointAddress, EndpointChannel}; use crate::client::lb::{ClusterDiscovery, XdsLbService}; @@ -367,17 +370,21 @@ impl XdsChannelBuilder { resource_manager: XdsResourceManager, ) -> XdsChannelGrpc { let router: Arc = Arc::new(XdsRouter::new(&cache)); + let circuit_breakers = ClusterCircuitBreakerRegistry::default(); #[cfg(feature = "_tls-any")] let discovery: Arc< dyn ClusterDiscovery>, > = Arc::new(XdsClusterDiscovery::new( - cache, + cache.clone(), GrpcMakeConnector::new(cert_provider_registry), )); #[cfg(not(feature = "_tls-any"))] let discovery: Arc< dyn ClusterDiscovery>, - > = Arc::new(XdsClusterDiscovery::new(cache, GrpcMakeConnector::new())); + > = Arc::new(XdsClusterDiscovery::new( + cache.clone(), + GrpcMakeConnector::new(), + )); let retry_policy = GrpcRetryPolicy::default(); let resources = Arc::new(XdsChannelResources { @@ -386,6 +393,10 @@ impl XdsChannelBuilder { }); let routing_layer = XdsRoutingLayer::new(router, self.pre_route.clone(), self.authority()); + let discovery = Arc::new( + CircuitBreakingClusterDiscovery::new(discovery, circuit_breakers) + .with_cluster_cache(cache), + ); let retry_layer = RetryLayer::new(retry_policy); let cluster_registry = Arc::new(ClusterClientRegistryGrpc::new()); let lb_service = XdsLbService::new(cluster_registry, discovery); @@ -420,8 +431,30 @@ impl XdsChannelBuilder { discovery: Arc>>, retry_policy: GrpcRetryPolicy, interceptor: Option>, + ) -> XdsChannelGrpc { + self.build_grpc_channel_from_parts_with_circuit_breakers( + router, + discovery, + retry_policy, + interceptor, + ClusterCircuitBreakerRegistry::new_for_test(), + ) + } + + #[cfg(test)] + pub(crate) fn build_grpc_channel_from_parts_with_circuit_breakers( + &self, + router: Arc, + discovery: Arc>>, + retry_policy: GrpcRetryPolicy, + interceptor: Option>, + circuit_breakers: ClusterCircuitBreakerRegistry, ) -> XdsChannelGrpc { let routing_layer = XdsRoutingLayer::new(router, interceptor, self.authority()); + let discovery = Arc::new(CircuitBreakingClusterDiscovery::new( + discovery, + circuit_breakers, + )); let retry_layer = RetryLayer::new(retry_policy); let cluster_registry = Arc::new(ClusterClientRegistryGrpc::new()); let lb_service = XdsLbService::new(cluster_registry, discovery); @@ -451,6 +484,7 @@ mod tests { use super::{XdsChannelBuilder, XdsChannelConfig}; use crate::XdsUri; use crate::client::channel::XdsChannelGrpc; + use crate::client::circuit_breaking::ClusterCircuitBreakerRegistry; use crate::client::endpoint::EndpointAddress; use crate::client::endpoint::EndpointChannel; @@ -468,6 +502,7 @@ mod tests { use crate::testutil::grpc::TestServer; use crate::xds::cache::XdsCache; use crate::xds::resource::EndpointsResource; + use crate::xds::resource::circuit_breaking::CircuitBreakingConfig; use crate::xds::resource::route_config::RouteConfigResource; use std::sync::Arc; use tokio::sync::mpsc; @@ -702,16 +737,223 @@ mod tests { assert_eq!(response.into_inner().message, "retry-server: retry-test"); } + #[tokio::test] + async fn test_xds_channel_enforces_injected_circuit_breaking_limit() { + use crate::client::retry::{GrpcRetryClassifier, GrpcRetryPolicy, RetryConfig}; + + let (_, servers) = setup_grpc_servers(1).await; + let xds_manager = Arc::new(MockXdsManager::from_test_servers(&servers)); + let circuit_breakers = ClusterCircuitBreakerRegistry::new_for_test(); + circuit_breakers.set_config("test-cluster", CircuitBreakingConfig { max_requests: 0 }); + + let retry_policy = GrpcRetryPolicy::new( + RetryConfig::new().num_retries(1), + GrpcRetryClassifier { + retry_on: vec![tonic::Code::Unavailable], + }, + ); + let xds_channel = XdsChannelBuilder::new(test_config()) + .build_grpc_channel_from_parts_with_circuit_breakers( + xds_manager.clone(), + xds_manager.clone(), + retry_policy, + None, + circuit_breakers, + ); + let mut client = GreeterClient::new(xds_channel); + + let error = client + .say_hello(HelloRequest { + name: "limited".to_string(), + }) + .await + .unwrap_err(); + + assert_eq!(error.code(), tonic::Code::Unavailable); + assert!(error.message().contains("max_requests limit 0")); + + for server in servers { + let _ = server.shutdown.send(()); + let _ = server.handle.await; + } + } + + #[tokio::test] + async fn test_xds_channel_uses_cds_circuit_breaking_config() { + let cluster_name = "test-cluster"; + let (_, servers) = setup_grpc_servers(1).await; + + let cache = Arc::new(XdsCache::new()); + cache.update_route_config(make_test_route_config(cluster_name)); + cache.update_cluster( + cluster_name, + make_test_cluster_with_circuit_breaking( + cluster_name, + CircuitBreakingConfig { max_requests: 0 }, + ), + ); + cache.update_endpoints(cluster_name, make_test_endpoints(cluster_name, &servers)); + + let channel = build_xds_channel_from_cache(cache).await; + let mut client = GreeterClient::new(channel); + let error = client + .say_hello(HelloRequest { + name: "cds-limited".to_string(), + }) + .await + .unwrap_err(); + + assert_eq!(error.code(), tonic::Code::Unavailable); + assert!(error.message().contains("max_requests limit 0")); + + for server in servers { + let _ = server.shutdown.send(()); + let _ = server.handle.await; + } + } + + #[tokio::test] + async fn test_xds_channel_waits_for_cds_before_circuit_breaking() { + let cluster_name = "test-cluster"; + let (_, servers) = setup_grpc_servers(1).await; + + let cache = Arc::new(XdsCache::new()); + cache.update_route_config(make_test_route_config(cluster_name)); + cache.update_endpoints(cluster_name, make_test_endpoints(cluster_name, &servers)); + + let channel = build_xds_channel_from_cache(cache.clone()).await; + let mut client = GreeterClient::new(channel); + let mut request = Box::pin(client.say_hello(HelloRequest { + name: "wait-for-cds".to_string(), + })); + + let early = + tokio::time::timeout(tokio::time::Duration::from_millis(20), request.as_mut()).await; + assert!( + early.is_err(), + "request should wait for CDS before acquiring a circuit-breaking permit", + ); + + cache.update_cluster( + cluster_name, + make_test_cluster_with_circuit_breaking( + cluster_name, + CircuitBreakingConfig { max_requests: 0 }, + ), + ); + + let result = tokio::time::timeout(tokio::time::Duration::from_secs(2), request) + .await + .expect("request should complete after CDS update"); + let error = result.unwrap_err(); + assert_eq!(error.code(), tonic::Code::Unavailable); + assert!(error.message().contains("max_requests limit 0")); + + for server in servers { + let _ = server.shutdown.send(()); + let _ = server.handle.await; + } + } + + #[tokio::test] + async fn test_xds_channel_recovers_cluster_after_remove_and_readd() { + let cluster_name = "test-cluster-remove-readd"; + let (_, servers) = setup_grpc_servers(2).await; + + let cache = Arc::new(XdsCache::new()); + cache.update_route_config(make_test_route_config(cluster_name)); + cache.update_cluster( + cluster_name, + make_test_cluster_with_circuit_breaking( + cluster_name, + CircuitBreakingConfig { max_requests: 1 }, + ), + ); + cache.update_endpoints( + cluster_name, + make_test_endpoints(cluster_name, &servers[..1]), + ); + + let channel = build_xds_channel_from_cache(cache.clone()).await; + let mut client = GreeterClient::new(channel); + let first = client + .say_hello(HelloRequest { + name: "before-removal".to_string(), + }) + .await + .unwrap(); + assert!(first.into_inner().message.starts_with("server-0:")); + + cache.remove_cluster(cluster_name); + cache.remove_endpoints(cluster_name); + cache.update_cluster( + cluster_name, + make_test_cluster_with_circuit_breaking( + cluster_name, + CircuitBreakingConfig { max_requests: 0 }, + ), + ); + cache.update_endpoints( + cluster_name, + make_test_endpoints(cluster_name, &servers[1..]), + ); + + let limited = tokio::time::timeout( + tokio::time::Duration::from_secs(2), + client.say_hello(HelloRequest { + name: "readded-limited".to_string(), + }), + ) + .await + .expect("re-added cluster should become ready") + .unwrap_err(); + assert_eq!(limited.code(), tonic::Code::Unavailable); + assert!(limited.message().contains("max_requests limit 0")); + + cache.update_cluster( + cluster_name, + make_test_cluster_with_circuit_breaking( + cluster_name, + CircuitBreakingConfig { max_requests: 1 }, + ), + ); + tokio::task::yield_now().await; + + let second = tokio::time::timeout( + tokio::time::Duration::from_secs(2), + client.say_hello(HelloRequest { + name: "after-readd".to_string(), + }), + ) + .await + .expect("updated cluster should serve requests") + .unwrap(); + assert!(second.into_inner().message.starts_with("server-1:")); + + for server in servers { + let _ = server.shutdown.send(()); + let _ = server.handle.await; + } + } + /// Helper: creates a minimal plaintext `ClusterResource` for tests that /// drive `XdsClusterDiscovery`. The cluster watch in `discover_cluster` /// blocks until a cluster is in the cache. fn make_test_cluster(cluster_name: &str) -> Arc { + make_test_cluster_with_circuit_breaking(cluster_name, CircuitBreakingConfig::default()) + } + + fn make_test_cluster_with_circuit_breaking( + cluster_name: &str, + circuit_breaking: CircuitBreakingConfig, + ) -> Arc { use crate::xds::resource::cluster::{ClusterResource, LbPolicy}; Arc::new(ClusterResource { name: cluster_name.to_string(), eds_service_name: None, lb_policy: LbPolicy::RoundRobin, security: None, + circuit_breaking, }) } @@ -763,32 +1005,25 @@ mod tests { /// Builds an XdsChannelGrpc using real XdsRouter and XdsClusterDiscovery /// backed by the given cache. async fn build_xds_channel_from_cache(cache: Arc) -> XdsChannelGrpc { - use crate::xds::cluster_discovery::{GrpcMakeConnector, XdsClusterDiscovery}; - use crate::xds::routing::XdsRouter; - - let router: Arc = Arc::new(XdsRouter::new(&cache)); + use crate::xds::resource_manager::XdsResourceManager; + let xds_client = xds_client::XdsClient::disconnected(); + let resource_manager = + XdsResourceManager::new(xds_client.clone(), cache.clone(), "test-listener".into()); + let builder = XdsChannelBuilder::new(test_config()); #[cfg(feature = "_tls-any")] - let discovery: Arc< - dyn ClusterDiscovery>, - > = { + { use crate::xds::cert_provider::CertProviderRegistry; let registry = Arc::new( CertProviderRegistry::from_bootstrap(&Default::default(), Default::default()) .unwrap(), ); - Arc::new(XdsClusterDiscovery::new( - cache, - GrpcMakeConnector::new(registry), - )) - }; + builder.build_from_cache(cache, registry, xds_client, resource_manager) + } #[cfg(not(feature = "_tls-any"))] - let discovery: Arc< - dyn ClusterDiscovery>, - > = Arc::new(XdsClusterDiscovery::new(cache, GrpcMakeConnector::new())); - - let builder = XdsChannelBuilder::new(test_config()); - builder.build_grpc_channel_from_parts(router, discovery, GrpcRetryPolicy::default(), None) + { + builder.build_from_cache(cache, xds_client, resource_manager) + } } /// Tests the full xDS stack (XdsRouter + XdsClusterDiscovery) with a diff --git a/tonic-xds/src/client/circuit_breaking.rs b/tonic-xds/src/client/circuit_breaking.rs index 1805acabe..77f9c6e29 100644 --- a/tonic-xds/src/client/circuit_breaking.rs +++ b/tonic-xds/src/client/circuit_breaking.rs @@ -34,15 +34,23 @@ use std::task::{Context, Poll}; use arc_swap::ArcSwapOption; use bytes::Bytes; use dashmap::DashMap; +use futures_util::StreamExt as _; use http::{Request, Response}; use http_body::{Body, Frame}; use pin_project_lite::pin_project; +use tokio::sync::watch; use tonic::body::Body as TonicBody; -use tower::{BoxError, Layer, Service}; +use tower::Layer; +use tower::discover::Change; +use tower::load::Load; +use tower::{BoxError, Service}; +use crate::client::lb::{BoxDiscover, ClusterDiscovery}; use crate::client::route::RouteDecision; -use crate::common::async_util::BoxFuture; -use crate::xds::resource::circuit_breaking::CircuitBreakingConfig; +use crate::common::async_util::{AbortOnDrop, BoxFuture}; +use crate::xds::cache::{CacheEvent, XdsCache}; +use crate::xds::resource::ClusterResource; +use crate::xds::resource::circuit_breaking::{CircuitBreakingConfig, DEFAULT_MAX_REQUESTS}; static GLOBAL_COUNTERS: OnceLock> = OnceLock::new(); @@ -85,6 +93,12 @@ impl ClusterCircuitBreakerRegistry { } } + #[allow(dead_code)] + pub(crate) fn set_config(&self, cluster: impl Into, config: CircuitBreakingConfig) { + let cluster = cluster.into(); + self.set_cluster_config(cluster.clone(), cluster, config); + } + pub(crate) fn set_cluster_config( &self, cluster: impl Into, @@ -103,9 +117,69 @@ impl ClusterCircuitBreakerRegistry { counter_key, counter, }, + 0, ); } + fn ensure_cluster_watch( + &self, + cache: &Arc, + cluster: &str, + state: &Arc, + ) -> u64 { + let mut lifecycle = state + .watch_lifecycle + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if lifecycle.running { + return lifecycle.generation; + } + + let mut cluster_watch = cache.watch_cluster(cluster); + lifecycle.generation = lifecycle.generation.wrapping_add(1).max(1); + let generation = lifecycle.generation; + lifecycle.running = true; + state + .active_watch_generation + .store(generation, Ordering::Release); + + let weak_registry = Arc::downgrade(&self.inner); + let weak_state = Arc::downgrade(state); + let task = tokio::spawn(async move { + while let Some(event) = cluster_watch.next_event().await { + let CacheEvent::Resource { + resource: cluster_resource, + .. + } = event + else { + break; + }; + let Some(inner) = weak_registry.upgrade() else { + return; + }; + let Some(state) = weak_state.upgrade() else { + return; + }; + let circuit_breakers = ClusterCircuitBreakerRegistry { inner }; + let config = CircuitBreakerRuntimeConfig::from_cluster( + &cluster_resource, + &circuit_breakers.inner.counters, + ); + circuit_breakers.update_state_config(&state, config, generation); + } + + let Some(inner) = weak_registry.upgrade() else { + return; + }; + let Some(state) = weak_state.upgrade() else { + return; + }; + ClusterCircuitBreakerRegistry { inner }.finish_cluster_watch(&state, generation); + }); + lifecycle._task = Some(AbortOnDrop(task)); + generation + } + fn ensure_state(&self, cluster: &str) -> Arc { if let Some(state) = self.inner.configs.get(cluster) { return state.clone(); @@ -119,25 +193,60 @@ impl ClusterCircuitBreakerRegistry { } fn cluster_breaker(&self, cluster: &str) -> Arc { + self.cluster_breaker_with_optional_cache(cluster, None) + } + + fn cluster_breaker_with_cache( + &self, + cluster: &str, + cluster_cache: Arc, + ) -> Arc { + self.cluster_breaker_with_optional_cache(cluster, Some(cluster_cache)) + } + + fn cluster_breaker_with_optional_cache( + &self, + cluster: &str, + cluster_cache: Option>, + ) -> Arc { let state = self.ensure_state(cluster); - Arc::new(ClusterCircuitBreaker { - cluster: Arc::from(cluster), + let cluster: Arc = Arc::from(cluster); + let default_counter_key = CounterKey::same_cluster(cluster.clone()); + let breaker = Arc::new(ClusterCircuitBreaker { + cluster, state, counters: self.inner.counters.clone(), - }) + default_counter_key, + registry: self.clone(), + cluster_cache, + }); + if let Some(cache) = breaker.cluster_cache.as_ref() { + self.ensure_cluster_watch(cache, &breaker.cluster, &breaker.state); + } + breaker } fn update_state_config( &self, state: &ClusterCircuitBreakerState, config: CircuitBreakerRuntimeConfig, + watch_generation: u64, ) { let _update_guard = state .update_lock .lock() .unwrap_or_else(|poisoned| poisoned.into_inner()); + if watch_generation != 0 + && state.active_watch_generation.load(Ordering::Acquire) != watch_generation + { + return; + } let previous = state.current_config(); if previous.as_deref() == Some(&config) { + state + .config_watch_generation + .store(watch_generation, Ordering::Release); + state.notify_config_changed(); return; } @@ -153,6 +262,27 @@ impl ClusterCircuitBreakerRegistry { if counter_key_changed && let Some(previous) = previous { self.deactivate_config(previous); } + state + .config_watch_generation + .store(watch_generation, Ordering::Release); + state.notify_config_changed(); + } + + fn finish_cluster_watch(&self, state: &Arc, generation: u64) { + let mut lifecycle = state + .watch_lifecycle + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()); + if lifecycle.generation != generation { + return; + } + + self.clear_state(state); + state + .finished_watch_generation + .store(generation, Ordering::Release); + lifecycle.running = false; + state.notify_config_changed(); } fn clear_state(&self, state: &ClusterCircuitBreakerState) { @@ -163,6 +293,7 @@ impl ClusterCircuitBreakerRegistry { if let Some(previous) = state.config.swap(None) { self.deactivate_config(previous); } + state.config_watch_generation.store(0, Ordering::Release); } fn deactivate_config(&self, config: Arc) { @@ -172,18 +303,16 @@ impl ClusterCircuitBreakerRegistry { self.inner.counters.cleanup_if_unused(&counter_key); } - fn acquire(&self, cluster: &str) -> Result, CircuitBreakerLimit> { + fn acquire(&self, cluster: &str) -> Result { self.cluster_breaker(cluster).acquire() } #[cfg(test)] fn in_flight(&self, cluster: &str) -> u32 { let breaker = self.cluster_breaker(cluster); - breaker - .state - .current_config() - .map(|config| self.inner.counters.in_flight(&config.counter_key)) - .unwrap_or(0) + self.inner + .counters + .in_flight(&breaker.current_counter_key()) } #[cfg(test)] @@ -236,6 +365,18 @@ impl PartialEq for CircuitBreakerRuntimeConfig { impl Eq for CircuitBreakerRuntimeConfig {} +impl CircuitBreakerRuntimeConfig { + fn from_cluster(cluster: &ClusterResource, counters: &ClusterRequestCounters) -> Self { + let counter_key = CounterKey::new(cluster.name.as_str(), cluster.eds_service_name()); + let counter = counters.counter(&counter_key); + Self { + max_requests: cluster.circuit_breaking.max_requests, + counter_key, + counter, + } + } +} + #[derive(Clone, Debug, Eq, Hash, PartialEq)] struct CounterKey { cluster: Arc, @@ -249,32 +390,75 @@ impl CounterKey { eds_service_name: eds_service_name.into(), } } + + fn same_cluster(cluster: Arc) -> Self { + Self { + cluster: cluster.clone(), + eds_service_name: cluster, + } + } } struct ClusterCircuitBreakerState { config: ArcSwapOption, update_lock: Mutex<()>, dropped_requests: AtomicU64, + watch_lifecycle: Mutex, + active_watch_generation: AtomicU64, + config_watch_generation: AtomicU64, + finished_watch_generation: AtomicU64, + config_version: watch::Sender, +} + +#[derive(Default)] +struct ClusterWatchLifecycle { + generation: u64, + running: bool, + _task: Option, } impl fmt::Debug for ClusterCircuitBreakerState { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let watch_running = self + .watch_lifecycle + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) + .running; f.debug_struct("ClusterCircuitBreakerState") .field("current_config", &self.current_config()) .field( "dropped_requests", &self.dropped_requests.load(Ordering::Acquire), ) + .field("watch_running", &watch_running) + .field( + "active_watch_generation", + &self.active_watch_generation.load(Ordering::Acquire), + ) + .field( + "config_watch_generation", + &self.config_watch_generation.load(Ordering::Acquire), + ) + .field( + "finished_watch_generation", + &self.finished_watch_generation.load(Ordering::Acquire), + ) .finish() } } impl ClusterCircuitBreakerState { fn new() -> Self { + let (config_version, _) = watch::channel(0); Self { config: ArcSwapOption::empty(), update_lock: Mutex::new(()), dropped_requests: AtomicU64::new(0), + watch_lifecycle: Mutex::new(ClusterWatchLifecycle::default()), + active_watch_generation: AtomicU64::new(0), + config_watch_generation: AtomicU64::new(0), + finished_watch_generation: AtomicU64::new(0), + config_version, } } @@ -286,23 +470,76 @@ impl ClusterCircuitBreakerState { self.dropped_requests.fetch_add(1, Ordering::AcqRel); } + fn notify_config_changed(&self) { + self.config_version + .send_modify(|version| *version = version.wrapping_add(1)); + } + #[cfg(test)] fn dropped_requests(&self) -> u64 { self.dropped_requests.load(Ordering::Acquire) } } -#[derive(Debug)] struct ClusterCircuitBreaker { cluster: Arc, state: Arc, counters: ClusterRequestCounters, + default_counter_key: CounterKey, + registry: ClusterCircuitBreakerRegistry, + cluster_cache: Option>, +} + +impl fmt::Debug for ClusterCircuitBreaker { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ClusterCircuitBreaker") + .field("cluster", &self.cluster) + .field("state", &self.state) + .field("watching_cluster_cache", &self.cluster_cache.is_some()) + .finish() + } } impl ClusterCircuitBreaker { - fn acquire(&self) -> Result, CircuitBreakerLimit> { - let Some(config) = self.state.current_config() else { - return Ok(None); + fn acquire(&self) -> Result { + if let Some(config) = self.state.current_config() { + return self.acquire_with_config( + config.counter_key.clone(), + config.counter.clone(), + config.max_requests, + ); + } + + let counter = self.counters.counter(&self.default_counter_key); + self.acquire_with_config( + self.default_counter_key.clone(), + counter, + DEFAULT_MAX_REQUESTS, + ) + } + + async fn acquire_when_ready(&self) -> Result> { + let Some(cache) = self.cluster_cache.as_ref() else { + return self + .acquire() + .map_err(|limit| limit_exceeded_response(&self.cluster, limit)); + }; + + if let Some(config) = self.current_watched_config() { + return self + .acquire_with_config( + config.counter_key.clone(), + config.counter.clone(), + config.max_requests, + ) + .map_err(|limit| limit_exceeded_response(&self.cluster, limit)); + } + + let generation = self + .registry + .ensure_cluster_watch(cache, &self.cluster, &self.state); + let Some(config) = self.wait_for_config(generation).await else { + return Err(cluster_unavailable_response(&self.cluster)); }; self.acquire_with_config( @@ -310,7 +547,49 @@ impl ClusterCircuitBreaker { config.counter.clone(), config.max_requests, ) - .map(Some) + .map_err(|limit| limit_exceeded_response(&self.cluster, limit)) + } + + fn current_watched_config(&self) -> Option> { + let generation = self.state.active_watch_generation.load(Ordering::Acquire); + if generation == 0 { + return None; + } + self.watched_config_for_generation(generation) + } + + fn watched_config_for_generation( + &self, + generation: u64, + ) -> Option> { + if self.state.config_watch_generation.load(Ordering::Acquire) != generation { + return None; + } + let config = self.state.current_config(); + if self.state.active_watch_generation.load(Ordering::Acquire) == generation + && self.state.config_watch_generation.load(Ordering::Acquire) == generation + { + config + } else { + None + } + } + + async fn wait_for_config(&self, generation: u64) -> Option> { + let mut config_version = self.state.config_version.subscribe(); + loop { + if let Some(config) = self.watched_config_for_generation(generation) { + return Some(config); + } + if self.state.finished_watch_generation.load(Ordering::Acquire) == generation + || self.state.active_watch_generation.load(Ordering::Acquire) != generation + { + return None; + } + if config_version.changed().await.is_err() { + return None; + } + } } fn acquire_with_config( @@ -328,6 +607,13 @@ impl ClusterCircuitBreaker { } } } + + fn current_counter_key(&self) -> CounterKey { + self.state + .current_config() + .map(|config| config.counter_key.clone()) + .unwrap_or_else(|| self.default_counter_key.clone()) + } } #[derive(Clone, Debug)] @@ -465,7 +751,7 @@ impl InFlightCounter { } #[derive(Debug)] -struct CircuitBreakerPermit { +pub(crate) struct CircuitBreakerPermit { counter: Option>, counter_key: CounterKey, counters: ClusterRequestCounters, @@ -484,6 +770,141 @@ impl Drop for CircuitBreakerPermit { } } +#[derive(Clone)] +pub(crate) struct CircuitBreakingClusterDiscovery { + inner: Arc>, + circuit_breakers: ClusterCircuitBreakerRegistry, + cluster_cache: Option>, +} + +impl CircuitBreakingClusterDiscovery { + pub(crate) fn new( + inner: Arc>, + circuit_breakers: ClusterCircuitBreakerRegistry, + ) -> Self { + Self { + inner, + circuit_breakers, + cluster_cache: None, + } + } + + pub(crate) fn with_cluster_cache(mut self, cache: Arc) -> Self { + self.cluster_cache = Some(cache); + self + } +} + +impl fmt::Debug for CircuitBreakingClusterDiscovery { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CircuitBreakingClusterDiscovery") + .field("circuit_breakers", &self.circuit_breakers) + .field("watching_cluster_cache", &self.cluster_cache.is_some()) + .finish() + } +} + +impl ClusterDiscovery> + for CircuitBreakingClusterDiscovery +where + Endpoint: Send + 'static, + S: Send + 'static, +{ + fn discover_cluster( + &self, + cluster_name: &str, + ) -> BoxDiscover> { + let breaker = match self.cluster_cache.clone() { + Some(cache) => self + .circuit_breakers + .cluster_breaker_with_cache(cluster_name, cache), + None => self.circuit_breakers.cluster_breaker(cluster_name), + }; + Box::pin( + self.inner + .discover_cluster(cluster_name) + .map(move |change| { + change.map(|change| match change { + Change::Insert(endpoint, service) => Change::Insert( + endpoint, + CircuitBreakingEndpointService::new(service, breaker.clone()), + ), + Change::Remove(endpoint) => Change::Remove(endpoint), + }) + }), + ) + } +} + +#[derive(Clone)] +pub(crate) struct CircuitBreakingEndpointService { + inner: S, + breaker: Arc, +} + +impl CircuitBreakingEndpointService { + fn new(inner: S, breaker: Arc) -> Self { + Self { inner, breaker } + } +} + +impl fmt::Debug for CircuitBreakingEndpointService { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CircuitBreakingEndpointService") + .field("inner", &self.inner) + .field("breaker", &self.breaker) + .finish() + } +} + +impl Service> for CircuitBreakingEndpointService +where + S: Service, Response = Response, Error: Into> + + Clone + + Send + + 'static, + S::Future: Send + 'static, + B: Send + 'static, +{ + type Response = Response; + type Error = BoxError; + type Future = BoxFuture>; + + fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.poll_ready(cx).map_err(Into::into) + } + + fn call(&mut self, request: Request) -> Self::Future { + let breaker = self.breaker.clone(); + let clone = self.inner.clone(); + let mut inner = std::mem::replace(&mut self.inner, clone); + Box::pin(async move { + let permit = match breaker.acquire_when_ready().await { + Ok(permit) => permit, + Err(response) => return Ok(response), + }; + + let response = inner.call(request).await.map_err(Into::into)?; + Ok(hold_permit(response, permit)) + }) + } +} + +impl Load for CircuitBreakingEndpointService { + type Metric = S::Metric; + + fn load(&self) -> Self::Metric { + self.inner.load() + } +} + +pub(crate) fn hold_permit( + response: Response, + permit: CircuitBreakerPermit, +) -> Response { + response.map(|body| TonicBody::new(PermitBody::new(body, permit))) +} + /// Tower layer that enforces A32 max in-flight requests per xDS cluster. /// /// This layer must wrap the ready per-cluster dispatch service inside retries so @@ -598,12 +1019,7 @@ where let mut inner = std::mem::replace(&mut self.inner, clone); Box::pin(async move { let response = inner.call(request).await.map_err(Into::into)?; - match permit { - Some(permit) => { - Ok(response.map(|body| TonicBody::new(PermitBody::new(body, permit)))) - } - None => Ok(response), - } + Ok(response.map(|body| TonicBody::new(PermitBody::new(body, permit)))) }) } } @@ -619,7 +1035,10 @@ enum CircuitBreakingError { /// The retry layer uses this extension to distinguish local `UNAVAILABLE` drops /// from retryable responses returned by an upstream service. #[derive(Clone, Copy, Debug)] -struct LocalCircuitBreakerDrop; +pub(crate) struct LocalCircuitBreakerDrop; + +#[derive(Clone, Copy, Debug)] +pub(crate) struct LocalPreEndpointResponse; pub(crate) fn is_local_circuit_breaker_drop(response: &Response) -> bool { response @@ -628,12 +1047,28 @@ pub(crate) fn is_local_circuit_breaker_drop(response: &Response) -> bool { .is_some() } +pub(crate) fn is_local_pre_endpoint_response(response: &Response) -> bool { + response + .extensions() + .get::() + .is_some() +} + fn limit_exceeded_response(cluster: &str, limit: CircuitBreakerLimit) -> Response { let mut response = status_response(tonic::Status::unavailable(format!( "circuit breaker open for cluster '{cluster}': max_requests limit {} reached", limit.max_requests, ))); response.extensions_mut().insert(LocalCircuitBreakerDrop); + response.extensions_mut().insert(LocalPreEndpointResponse); + response +} + +fn cluster_unavailable_response(cluster: &str) -> Response { + let mut response = status_response(tonic::Status::unavailable(format!( + "cluster '{cluster}' is no longer available", + ))); + response.extensions_mut().insert(LocalPreEndpointResponse); response } @@ -724,7 +1159,6 @@ mod tests { use super::*; const CLUSTER: &str = "cluster-a"; - const EDS_SERVICE_NAME: &str = "eds-service-a"; fn request() -> Request { let mut request = Request::new(TonicBody::empty()); @@ -737,14 +1171,20 @@ mod tests { fn configured_breakers(max_requests: u32) -> ClusterCircuitBreakerRegistry { let breakers = ClusterCircuitBreakerRegistry::new_for_test(); - breakers.set_cluster_config( - CLUSTER, - EDS_SERVICE_NAME, - CircuitBreakingConfig { max_requests }, - ); + breakers.set_config(CLUSTER, CircuitBreakingConfig { max_requests }); breakers } + fn cluster_resource(max_requests: u32) -> Arc { + Arc::new(ClusterResource { + name: CLUSTER.to_string(), + eds_service_name: None, + lb_policy: crate::xds::resource::cluster::LbPolicy::RoundRobin, + security: None, + circuit_breaking: CircuitBreakingConfig { max_requests }, + }) + } + #[tokio::test] async fn rejects_requests_when_cluster_limit_is_reached() { let breakers = configured_breakers(1); @@ -818,11 +1258,7 @@ mod tests { .unwrap(); assert_eq!(breakers.in_flight(CLUSTER), 2); - breakers.set_cluster_config( - CLUSTER, - EDS_SERVICE_NAME, - CircuitBreakingConfig { max_requests: 1 }, - ); + breakers.set_config(CLUSTER, CircuitBreakingConfig { max_requests: 1 }); let third = service .ready() @@ -1021,24 +1457,16 @@ mod tests { let counters = ClusterRequestCounters::isolated(); let first_breakers = ClusterCircuitBreakerRegistry::with_counters(counters.clone()); let second_breakers = ClusterCircuitBreakerRegistry::with_counters(counters.clone()); - first_breakers.set_cluster_config( - CLUSTER, - EDS_SERVICE_NAME, - CircuitBreakingConfig { max_requests: 1 }, - ); - second_breakers.set_cluster_config( - CLUSTER, - EDS_SERVICE_NAME, - CircuitBreakingConfig { max_requests: 1 }, - ); + first_breakers.set_config(CLUSTER, CircuitBreakingConfig { max_requests: 1 }); + second_breakers.set_config(CLUSTER, CircuitBreakingConfig { max_requests: 1 }); - let first = first_breakers.acquire(CLUSTER).unwrap().unwrap(); + let first = first_breakers.acquire(CLUSTER).unwrap(); assert!(second_breakers.acquire(CLUSTER).is_err()); assert_eq!(first_breakers.dropped_requests(CLUSTER), 0); assert_eq!(second_breakers.dropped_requests(CLUSTER), 1); drop(first); - let second = second_breakers.acquire(CLUSTER).unwrap().unwrap(); + let second = second_breakers.acquire(CLUSTER).unwrap(); drop(second); drop(first_breakers); drop(second_breakers); @@ -1046,28 +1474,89 @@ mod tests { } #[test] - fn unconfigured_cluster_has_no_limit_or_counter() { + fn default_limit_rejects_the_1025th_request() { let breakers = ClusterCircuitBreakerRegistry::new_for_test(); let breaker = breakers.cluster_breaker(CLUSTER); + let permits: Vec<_> = (0..DEFAULT_MAX_REQUESTS) + .map(|_| breaker.acquire().unwrap()) + .collect(); - for _ in 0..2048 { - assert!(breaker.acquire().unwrap().is_none()); - } + assert!(breaker.acquire().is_err()); + assert_eq!(breakers.dropped_requests(CLUSTER), 1); - assert_eq!(breakers.dropped_requests(CLUSTER), 0); + drop(permits); assert_eq!(breakers.counter_count(), 0); } + #[tokio::test] + async fn endpoint_waiting_for_ready_does_not_hold_permit() { + let breakers = configured_breakers(1); + let service = BackpressuredService { + ready_budget: Arc::new(AtomicU32::new(0)), + calls: Arc::new(AtomicU32::new(0)), + }; + let mut service = + CircuitBreakingEndpointService::new(service, breakers.cluster_breaker(CLUSTER)); + + let early = + tokio::time::timeout(tokio::time::Duration::from_millis(20), service.ready()).await; + + assert!( + early.is_err(), + "endpoint wrapper should wait for inner endpoint readiness", + ); + assert_eq!(breakers.in_flight(CLUSTER), 0); + } + + #[tokio::test] + async fn endpoint_limit_responses_are_not_retried() { + use crate::client::retry::{GrpcRetryClassifier, GrpcRetryPolicy, RetryConfig}; + + let breakers = configured_breakers(0); + let calls = Arc::new(AtomicU32::new(0)); + let call_counter = calls.clone(); + let service = service_fn( + move |_request: Request>| { + call_counter.fetch_add(1, Ordering::SeqCst); + async { Ok::<_, BoxError>(Response::new(TonicBody::new(PendingBody))) } + }, + ); + let service = + CircuitBreakingEndpointService::new(service, breakers.cluster_breaker(CLUSTER)); + let retry_policy = GrpcRetryPolicy::new( + RetryConfig::new().num_retries(1), + GrpcRetryClassifier { + retry_on: vec![tonic::Code::Unavailable], + }, + ); + let mut service = tower::ServiceBuilder::new() + .layer(RetryLayer::new(retry_policy)) + .service(service); + + let response = service + .ready() + .await + .unwrap() + .call(Request::new(TonicBody::empty())) + .await + .unwrap(); + let status = tonic::Status::from_header_map(response.headers()).unwrap(); + + assert_eq!(status.code(), Code::Unavailable); + assert_eq!(calls.load(Ordering::SeqCst), 0); + assert_eq!(breakers.dropped_requests(CLUSTER), 1); + } + #[test] fn eds_service_name_change_uses_independent_counter() { let breakers = ClusterCircuitBreakerRegistry::new_for_test(); breakers.set_cluster_config(CLUSTER, "eds-a", CircuitBreakingConfig { max_requests: 1 }); - let first = breakers.acquire(CLUSTER).unwrap().unwrap(); + let first = breakers.acquire(CLUSTER).unwrap(); assert_eq!(breakers.in_flight(CLUSTER), 1); breakers.set_cluster_config(CLUSTER, "eds-b", CircuitBreakingConfig { max_requests: 1 }); assert_eq!(breakers.in_flight(CLUSTER), 0); - let second = breakers.acquire(CLUSTER).unwrap().unwrap(); + let second = breakers.acquire(CLUSTER).unwrap(); assert_eq!(breakers.in_flight(CLUSTER), 1); assert!(breakers.acquire(CLUSTER).is_err()); @@ -1094,7 +1583,7 @@ mod tests { #[test] fn cluster_removal_cleans_up_counter_after_in_flight_requests_finish() { let breakers = configured_breakers(1); - let permit = breakers.acquire(CLUSTER).unwrap().unwrap(); + let permit = breakers.acquire(CLUSTER).unwrap(); assert_eq!(breakers.counter_count(), 1); breakers.clear_cluster_config(CLUSTER); @@ -1106,14 +1595,134 @@ mod tests { #[tokio::test] async fn cached_breaker_observes_cluster_removal_and_recreation() { + let cache = Arc::new(XdsCache::new()); + cache.update_cluster(CLUSTER, cluster_resource(1)); + let breakers = ClusterCircuitBreakerRegistry::new_for_test(); + let breaker = breakers.cluster_breaker_with_cache(CLUSTER, cache.clone()); + + let first = breaker.acquire_when_ready().await.unwrap(); + drop(first); + + let generation = breaker + .state + .active_watch_generation + .load(Ordering::Acquire); + cache.remove_cluster(CLUSTER); + tokio::time::timeout(Duration::from_secs(1), async { + while breaker + .state + .finished_watch_generation + .load(Ordering::Acquire) + != generation + { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + cache.update_cluster(CLUSTER, cluster_resource(1)); + let second = tokio::time::timeout(Duration::from_secs(1), breaker.acquire_when_ready()) + .await + .expect("re-added cluster should wake the cached breaker") + .expect("re-added cluster should be available"); + drop(second); + } + + #[tokio::test] + async fn removal_before_watch_task_runs_wakes_waiting_request() { + let cache = Arc::new(XdsCache::new()); + let breakers = ClusterCircuitBreakerRegistry::new_for_test(); + let breaker = breakers.cluster_breaker_with_cache(CLUSTER, cache.clone()); + + cache.remove_cluster(CLUSTER); + let response = tokio::time::timeout(Duration::from_secs(1), breaker.acquire_when_ready()) + .await + .expect("cluster removal should wake the request") + .expect_err("removed cluster should be unavailable"); + let status = tonic::Status::from_header_map(response.headers()).unwrap(); + assert_eq!(status.code(), Code::Unavailable); + assert!(is_local_pre_endpoint_response(&response)); + assert!(!is_local_circuit_breaker_drop(&response)); + } + + #[tokio::test] + async fn waiting_request_does_not_cross_watch_generations() { + let cache = Arc::new(XdsCache::new()); + let breakers = ClusterCircuitBreakerRegistry::new_for_test(); + let breaker = breakers.cluster_breaker_with_cache(CLUSTER, cache.clone()); + let old_generation = breaker + .state + .active_watch_generation + .load(Ordering::Acquire); + let mut old_request = Box::pin(breaker.acquire_when_ready()); + + std::future::poll_fn(|cx| { + assert!(matches!(old_request.as_mut().poll(cx), Poll::Pending)); + Poll::Ready(()) + }) + .await; + + cache.remove_cluster(CLUSTER); + cache.update_cluster(CLUSTER, cluster_resource(1)); + tokio::time::timeout(Duration::from_secs(1), async { + while breaker + .state + .finished_watch_generation + .load(Ordering::Acquire) + != old_generation + { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + + let new_permit = breaker.acquire_when_ready().await.unwrap(); + drop(new_permit); + + let response = old_request + .await + .expect_err("old request must not consume the re-added cluster config"); + assert_eq!( + tonic::Status::from_header_map(response.headers()) + .unwrap() + .code(), + Code::Unavailable, + ); + } + + #[tokio::test] + async fn dropping_cache_backed_breaker_releases_watcher_owners() { + let cache = Arc::new(XdsCache::new()); + cache.update_cluster(CLUSTER, cluster_resource(1)); + let breakers = ClusterCircuitBreakerRegistry::new_for_test(); + let breaker = breakers.cluster_breaker_with_cache(CLUSTER, cache.clone()); + let permit = breaker.acquire_when_ready().await.unwrap(); + drop(permit); + + let weak_cache = Arc::downgrade(&cache); + let weak_registry = Arc::downgrade(&breakers.inner); + drop(breaker); + drop(breakers); + drop(cache); + tokio::task::yield_now().await; + + assert!(weak_cache.upgrade().is_none()); + assert!(weak_registry.upgrade().is_none()); + } + + #[tokio::test] + async fn endpoint_holds_positive_limit_until_response_body_drops() { let breakers = configured_breakers(1); let calls = Arc::new(AtomicU32::new(0)); let call_counter = calls.clone(); let service = service_fn(move |_request: Request| { call_counter.fetch_add(1, Ordering::SeqCst); - async { Ok::<_, BoxError>(Response::new(TonicBody::empty())) } + async { Ok::<_, BoxError>(Response::new(TonicBody::new(PendingBody))) } }); - let mut service = CircuitBreakingLayer::new(breakers.clone()).layer(service); + let mut service = + CircuitBreakingEndpointService::new(service, breakers.cluster_breaker(CLUSTER)); let first = service .ready() @@ -1122,9 +1731,8 @@ mod tests { .call(request()) .await .unwrap(); - drop(first); + assert_eq!(breakers.in_flight(CLUSTER), 1); - breakers.clear_cluster_config(CLUSTER); let second = service .ready() .await @@ -1132,14 +1740,13 @@ mod tests { .call(request()) .await .unwrap(); - assert!(tonic::Status::from_header_map(second.headers()).is_none()); - assert_eq!(calls.load(Ordering::SeqCst), 2); + let status = tonic::Status::from_header_map(second.headers()).unwrap(); + assert_eq!(status.code(), Code::Unavailable); + assert_eq!(calls.load(Ordering::SeqCst), 1); + + drop(first); + assert_eq!(breakers.in_flight(CLUSTER), 0); - breakers.set_cluster_config( - CLUSTER, - EDS_SERVICE_NAME, - CircuitBreakingConfig { max_requests: 0 }, - ); let third = service .ready() .await @@ -1147,21 +1754,16 @@ mod tests { .call(request()) .await .unwrap(); - let status = tonic::Status::from_header_map(third.headers()).unwrap(); - assert_eq!(status.code(), Code::Unavailable); assert_eq!(calls.load(Ordering::SeqCst), 2); + drop(third); } #[test] fn dropping_breakers_releases_config_counter_ref() { let counters = ClusterRequestCounters::isolated(); let breakers = ClusterCircuitBreakerRegistry::with_counters(counters.clone()); - breakers.set_cluster_config( - CLUSTER, - EDS_SERVICE_NAME, - CircuitBreakingConfig { max_requests: 1 }, - ); - let permit = breakers.acquire(CLUSTER).unwrap().unwrap(); + breakers.set_config(CLUSTER, CircuitBreakingConfig { max_requests: 1 }); + let permit = breakers.acquire(CLUSTER).unwrap(); drop(permit); assert_eq!(counters.counter_count(), 1); @@ -1173,7 +1775,7 @@ mod tests { #[test] fn cleanup_keeps_counter_with_outstanding_clone() { let counters = ClusterRequestCounters::isolated(); - let counter_key = CounterKey::new(CLUSTER, EDS_SERVICE_NAME); + let counter_key = CounterKey::new(CLUSTER, CLUSTER); let counter = counters.counter(&counter_key); counters.cleanup_if_unused(&counter_key); @@ -1197,6 +1799,19 @@ mod tests { assert_eq!(counters.counter_count(), 2); } + #[test] + fn cached_cluster_breaker_does_not_pin_default_counter() { + let counters = ClusterRequestCounters::isolated(); + let breakers = ClusterCircuitBreakerRegistry::with_counters(counters.clone()); + let breaker = breakers.cluster_breaker(CLUSTER); + let permit = breaker.acquire().unwrap(); + assert_eq!(counters.counter_count(), 1); + + drop(permit); + assert_eq!(counters.counter_count(), 0); + assert_eq!(breaker.cluster.as_ref(), CLUSTER); + } + #[tokio::test] async fn waiting_for_inner_readiness_does_not_acquire_permit() { let breakers = configured_breakers(1); diff --git a/tonic-xds/src/client/loadbalance/loadbalancer.rs b/tonic-xds/src/client/loadbalance/loadbalancer.rs index d58392746..ccda52e0f 100644 --- a/tonic-xds/src/client/loadbalance/loadbalancer.rs +++ b/tonic-xds/src/client/loadbalance/loadbalancer.rs @@ -62,6 +62,7 @@ use tower::discover::{Change, Discover}; use arc_swap::ArcSwap; +use crate::client::circuit_breaking::is_local_pre_endpoint_response; use crate::client::endpoint::{Connector, EndpointAddress}; use crate::client::loadbalance::channel_state::{ EjectionConfig, IdleChannel, ReadyChannel, UnejectedChannel, @@ -72,6 +73,26 @@ use crate::client::loadbalance::outlier_detection::{OutlierDetector, OutlierStat use crate::client::loadbalance::pickers::ChannelPicker; use crate::xds::resource::outlier_detection::OutlierDetectionConfig; +trait LbResponseOutcome { + fn is_local_pre_endpoint_response(&self) -> bool { + false + } +} + +impl LbResponseOutcome for http::Response { + fn is_local_pre_endpoint_response(&self) -> bool { + is_local_pre_endpoint_response(self) + } +} + +fn outlier_outcome(result: &Result) -> Option { + match result { + Ok(response) if response.is_local_pre_endpoint_response() => None, + Ok(_) => Some(true), + Err(_) => Some(false), + } +} + /// Future returned by [`LoadBalancer::call`]. Either resolves /// immediately with an [`LbError`] or drives the selected channel. pub(crate) enum LbFuture { @@ -340,7 +361,7 @@ where D::Error: Into, C: Connector + Send + Sync + 'static, C::Service: Service + Clone + Send + 'static, - >::Response: Send + 'static, + >::Response: LbResponseOutcome + Send + 'static, >::Error: Into, >::Future: Send + 'static, Req: Send + 'static, @@ -393,7 +414,9 @@ where .await .map_err(|e| LbError::LbChannelPollReadyError(e.into()))?; let result = svc.call(req).await; - svc.record_outcome(result.is_ok()); + if let Some(success) = outlier_outcome(&result) { + svc.record_outcome(success); + } result.map_err(|e| LbError::LbChannelCallError(e.into())) })) } @@ -402,6 +425,7 @@ where #[cfg(test)] mod tests { use super::*; + use crate::client::circuit_breaking::LocalPreEndpointResponse; use crate::client::endpoint::Connector; use crate::client::loadbalance::pickers::p2c::P2cPicker; use crate::common::async_util::BoxFuture; @@ -467,6 +491,8 @@ mod tests { } } + impl LbResponseOutcome for &'static str {} + // -- Mock connector -- /// A connector that returns a pending future until signaled via oneshot. @@ -611,6 +637,22 @@ mod tests { Ok(Change::Remove(addr(port))) } + #[test] + fn circuit_breaker_drops_are_not_endpoint_outlier_outcomes() { + let normal: Result, tower::BoxError> = Ok(http::Response::new(())); + assert_eq!(outlier_outcome(&normal), Some(true)); + + let mut dropped_response = http::Response::new(()); + dropped_response + .extensions_mut() + .insert(LocalPreEndpointResponse); + let dropped: Result, tower::BoxError> = Ok(dropped_response); + assert_eq!(outlier_outcome(&dropped), None); + + let error: Result, tower::BoxError> = Err("endpoint error".into()); + assert_eq!(outlier_outcome(&error), Some(false)); + } + /// A burst of inserts drained in one poll rebuilds the ring exactly once, /// with the full member set (gRFC A42 ring-hash relies on this). A later /// drain (the remove) is a distinct, single rebuild. diff --git a/tonic-xds/src/xds/cache.rs b/tonic-xds/src/xds/cache.rs index 12ae472d5..4f7bb1128 100644 --- a/tonic-xds/src/xds/cache.rs +++ b/tonic-xds/src/xds/cache.rs @@ -35,42 +35,183 @@ use tokio::sync::watch; use crate::xds::resource::{ClusterResource, EndpointsResource, RouteConfigResource}; -/// A wrapper around [`watch::Receiver`] that exposes only a single `next()` -/// method, preventing misuse of the raw watch API. +enum WatchValue { + Pending { + generation: u64, + revision: u64, + }, + Resource { + generation: u64, + revision: u64, + resource: Arc, + }, + Removed { + generation: u64, + revision: u64, + }, +} + +impl Clone for WatchValue { + fn clone(&self) -> Self { + match self { + Self::Pending { + generation, + revision, + } => Self::Pending { + generation: *generation, + revision: *revision, + }, + Self::Resource { + generation, + revision, + resource, + } => Self::Resource { + generation: *generation, + revision: *revision, + resource: Arc::clone(resource), + }, + Self::Removed { + generation, + revision, + } => Self::Removed { + generation: *generation, + revision: *revision, + }, + } + } +} + +impl WatchValue { + fn generation(&self) -> u64 { + match self { + Self::Pending { generation, .. } + | Self::Resource { generation, .. } + | Self::Removed { generation, .. } => *generation, + } + } + + fn revision(&self) -> u64 { + match self { + Self::Pending { revision, .. } + | Self::Resource { revision, .. } + | Self::Removed { revision, .. } => *revision, + } + } +} + +pub(crate) enum CacheEvent { + Resource { + generation: u64, + revision: u64, + resource: Arc, + }, + Removed, +} + +/// A wrapper around [`watch::Receiver`] that preserves resource-removal events. pub(crate) struct CacheWatch { - rx: watch::Receiver>>, + rx: watch::Receiver>, + generation: u64, + revision: u64, + pending_resource: Option<(u64, u64, Arc)>, } impl CacheWatch { - fn new(mut rx: watch::Receiver>>) -> Self { + fn new(mut rx: watch::Receiver>) -> Self { + let generation = rx.borrow().generation(); + let revision = rx.borrow().revision(); // Ensure late watchers see the existing value on first next(). rx.mark_changed(); - Self { rx } + Self { + rx, + generation, + revision, + pending_resource: None, + } } /// Waits for the next resource update and returns it. /// - /// Returns `None` if the sender was dropped (resource removed from cache). + /// Resource removals are skipped; use [`Self::next_event`] when they must be + /// observed. Returns `None` only when the cache is dropped. pub(crate) async fn next(&mut self) -> Option> { + while let Some(event) = self.next_event().await { + if let CacheEvent::Resource { resource, .. } = event { + return Some(resource); + } + } + None + } + + /// Waits for the next resource update or removal. + pub(crate) async fn next_event(&mut self) -> Option> { + if let Some((generation, revision, resource)) = self.pending_resource.take() { + return Some(CacheEvent::Resource { + generation, + revision, + resource, + }); + } + loop { if self.rx.changed().await.is_err() { return None; } - let val = self.rx.borrow_and_update().clone(); - if val.is_some() { - return val; + match self.rx.borrow_and_update().clone() { + WatchValue::Pending { + generation, + revision, + } => { + self.generation = generation; + self.revision = revision; + } + WatchValue::Resource { + generation, + revision, + resource, + } if generation != self.generation => { + self.generation = generation; + self.revision = revision; + self.pending_resource = Some((generation, revision, resource)); + return Some(CacheEvent::Removed); + } + WatchValue::Resource { + generation, + revision, + resource, + } => { + self.generation = generation; + self.revision = revision; + return Some(CacheEvent::Resource { + generation, + revision, + resource, + }); + } + WatchValue::Removed { + generation, + revision, + } => { + self.generation = generation; + self.revision = revision; + return Some(CacheEvent::Removed); + } } } } + + pub(crate) fn current_revision(&self) -> u64 { + self.rx.borrow().revision() + } } /// A keyed collection of [`watch`] channels for a single xDS resource type. /// /// Each entry is lazily created on first access (watch or update) and -/// starts with `None`. Writers call [`update`](Self::update) to set the value; -/// consumers call [`watch`](Self::watch) to receive changes. +/// starts pending. Removed entries remain as tombstones so late watchers do not +/// mistake a known removal for a resource that has not arrived yet. struct WatchMap { - inner: DashMap>>>, + inner: DashMap>>, } impl WatchMap { @@ -85,7 +226,13 @@ impl WatchMap { /// Lazily creates the watch channel if this is the first access for the key. fn update(&self, key: &str, value: Arc) { let tx = self.ensure(key); - tx.send_replace(Some(value)); + tx.send_modify(move |current| { + *current = WatchValue::Resource { + generation: current.generation(), + revision: current.revision().wrapping_add(1), + resource: value, + }; + }); } /// Watches resource changes for the given key. @@ -96,18 +243,28 @@ impl WatchMap { CacheWatch::new(tx.subscribe()) } - /// Removes the watch channel for the given key. - /// - /// Dropping the sender closes all watcher receivers. + /// Marks the resource removed and notifies current and future watchers. fn remove(&self, key: &str) { - self.inner.remove(key); + let tx = self.ensure(key); + tx.send_modify(|current| { + *current = WatchValue::Removed { + generation: current.generation().wrapping_add(1), + revision: current.revision().wrapping_add(1), + }; + }); } /// Returns the sender for `key`, creating the watch channel if needed. - fn ensure(&self, key: &str) -> watch::Sender>> { + fn ensure(&self, key: &str) -> watch::Sender> { self.inner .entry(key.to_string()) - .or_insert_with(|| watch::channel(None).0) + .or_insert_with(|| { + watch::channel(WatchValue::Pending { + generation: 0, + revision: 0, + }) + .0 + }) .value() .clone() } @@ -124,7 +281,7 @@ impl WatchMap { /// CDS readiness, then [`watch_endpoints`](Self::watch_endpoints) to track EDS updates. pub(crate) struct XdsCache { /// Active route configuration (from LDS inline or RDS). - route_config_tx: watch::Sender>>, + route_config_tx: watch::Sender>, /// Per-cluster CDS state with readiness gating. clusters: WatchMap, @@ -136,7 +293,10 @@ pub(crate) struct XdsCache { impl XdsCache { /// Creates a new empty cache with no resources. pub(crate) fn new() -> Self { - let (route_config_tx, _) = watch::channel(None); + let (route_config_tx, _) = watch::channel(WatchValue::Pending { + generation: 0, + revision: 0, + }); Self { route_config_tx, clusters: WatchMap::new(), @@ -146,7 +306,13 @@ impl XdsCache { /// Updates the active route configuration and notifies all watchers. pub(crate) fn update_route_config(&self, config: Arc) { - self.route_config_tx.send_replace(Some(config)); + self.route_config_tx.send_modify(move |current| { + *current = WatchValue::Resource { + generation: current.generation(), + revision: current.revision().wrapping_add(1), + resource: config, + }; + }); } /// Watches route configuration changes. @@ -215,11 +381,14 @@ mod tests { } fn make_cluster(name: &str, lb: LbPolicy) -> Arc { + use crate::xds::resource::circuit_breaking::CircuitBreakingConfig; + Arc::new(ClusterResource { name: name.to_string(), eds_service_name: None, lb_policy: lb, security: None, + circuit_breaking: CircuitBreakingConfig::default(), }) } @@ -279,14 +448,64 @@ mod tests { } #[tokio::test] - async fn cluster_remove_closes_watcher() { + async fn cluster_remove_notifies_watcher() { let cache = XdsCache::new(); let mut watch = cache.watch_cluster("c1"); cache.update_cluster("c1", make_cluster("c1", LbPolicy::RoundRobin)); watch.next().await; // consume cache.remove_cluster("c1"); - assert!(watch.next().await.is_none()); + assert!(matches!( + watch.next_event().await, + Some(CacheEvent::Removed) + )); + } + + #[tokio::test] + async fn late_cluster_watcher_observes_removal_and_readd() { + let cache = XdsCache::new(); + cache.remove_cluster("c1"); + + let mut event_watch = cache.watch_cluster("c1"); + assert!(matches!( + event_watch.next_event().await, + Some(CacheEvent::Removed) + )); + + let mut standard_watch = cache.watch_cluster("c1"); + assert!( + tokio::time::timeout(std::time::Duration::from_millis(20), standard_watch.next()) + .await + .is_err(), + "generic late watchers should await a future re-add", + ); + + cache.update_cluster("c1", make_cluster("c1", LbPolicy::LeastRequest)); + let cluster = standard_watch.next().await.unwrap(); + assert_eq!(cluster.lb_policy, LbPolicy::LeastRequest); + } + + #[tokio::test] + async fn watcher_preserves_coalesced_removal_before_readd() { + let cache = XdsCache::new(); + let mut watch = cache.watch_cluster("c1"); + cache.update_cluster("c1", make_cluster("c1", LbPolicy::RoundRobin)); + watch.next().await.unwrap(); + + cache.remove_cluster("c1"); + cache.update_cluster("c1", make_cluster("c1", LbPolicy::LeastRequest)); + + assert!(matches!( + watch.next_event().await, + Some(CacheEvent::Removed) + )); + let Some(CacheEvent::Resource { + resource: cluster, .. + }) = watch.next_event().await + else { + panic!("expected re-added cluster after removal"); + }; + assert_eq!(cluster.lb_policy, LbPolicy::LeastRequest); } #[tokio::test] @@ -314,14 +533,17 @@ mod tests { } #[tokio::test] - async fn remove_endpoints_closes_watchers() { + async fn remove_endpoints_notifies_watchers() { let cache = XdsCache::new(); let mut watch = cache.watch_endpoints("c1"); cache.update_endpoints("c1", make_endpoints("c1")); watch.next().await; // consume cache.remove_endpoints("c1"); - assert!(watch.next().await.is_none()); + assert!(matches!( + watch.next_event().await, + Some(CacheEvent::Removed) + )); } #[tokio::test] diff --git a/tonic-xds/src/xds/cluster_discovery.rs b/tonic-xds/src/xds/cluster_discovery.rs index caabda919..7bf3e53c0 100644 --- a/tonic-xds/src/xds/cluster_discovery.rs +++ b/tonic-xds/src/xds/cluster_discovery.rs @@ -29,8 +29,8 @@ //! //! 1. The cluster resource watch — produces a fresh [`Connector`] on each //! CDS update (e.g. when `transport_socket` changes). The connector is -//! held inside a [`ConnectorSwap`] so the diff loop reads the latest -//! snapshot per endpoint connection. +//! published with its cache generation so re-added endpoints cannot use +//! transport settings from the removed cluster generation. //! 2. The endpoint watch — produces `Change::Insert` / `Change::Remove` //! events forwarded to the LB layer. //! @@ -40,8 +40,7 @@ use std::sync::Arc; use std::time::Duration; -use arc_swap::ArcSwap; -use tokio::sync::mpsc; +use tokio::sync::{mpsc, watch}; use tokio_stream::StreamExt as _; use tokio_stream::wrappers::ReceiverStream; use tonic::transport::{Channel, Endpoint}; @@ -52,12 +51,12 @@ use crate::client::endpoint::{ }; use crate::client::lb::{BoxDiscover, ClusterDiscovery}; use crate::common::async_util::BoxFuture; -use crate::xds::cache::XdsCache; +use crate::xds::cache::{CacheEvent, XdsCache}; #[cfg(feature = "_tls-any")] use crate::xds::cert_provider::verifier::XdsServerCertVerifier; #[cfg(feature = "_tls-any")] use crate::xds::cert_provider::{CertProviderRegistry, CertificateProvider}; -use crate::xds::endpoint_manager::{ConnectorSwap, EndpointManager}; +use crate::xds::endpoint_manager::{ConnectorState, EndpointManager}; use crate::xds::resource::security::ClusterSecurityConfig; /// Buffer capacity for the discovery channel between the spawned task and @@ -123,12 +122,24 @@ impl ClusterDiscovery for XdsCl tokio::spawn(async move { let mut cluster_watch = cache.watch_cluster(&cluster_name); - let connector_swap: ConnectorSwap = loop { - let Some(cluster) = cluster_watch.next().await else { + let (generation, connector) = loop { + let event = tokio::select! { + _ = tx.closed() => return, + event = cluster_watch.next_event() => event, + }; + let Some(event) = event else { return; }; + let CacheEvent::Resource { + generation, + resource: cluster, + .. + } = event + else { + continue; + }; match make_connector.make_connector(ClusterConfig::from_resource(&cluster)) { - Ok(c) => break Arc::new(ArcSwap::from_pointee(c)), + Ok(connector) => break (generation, connector), Err(e) => tracing::warn!( cluster = %cluster_name, error = %e, @@ -137,24 +148,52 @@ impl ClusterDiscovery for XdsCl } }; - let manager = EndpointManager::new(Arc::clone(&connector_swap)); + let (connector_tx, connector_rx) = + watch::channel(Arc::new(ConnectorState::new(generation, connector))); + let manager = EndpointManager::new(connector_rx); let mut endpoints = manager.discover_endpoints(cache.watch_endpoints(&cluster_name)); loop { tokio::select! { + _ = tx.closed() => return, Some(change) = endpoints.next() => { if tx.send(change).await.is_err() { return; } } - Some(cluster) = cluster_watch.next() => { + event = cluster_watch.next_event() => { + let Some(event) = event else { + return; + }; + let CacheEvent::Resource { + generation, + resource: cluster, + .. + } = event + else { + continue; + }; match make_connector.make_connector(ClusterConfig::from_resource(&cluster)) { - Ok(new) => connector_swap.store(Arc::new(new)), - Err(e) => tracing::warn!( - cluster = %cluster_name, - error = %e, - "CDS update rejected; keeping previous connector", - ), + Ok(connector) => { + connector_tx.send_replace(Arc::new(ConnectorState::new( + generation, + connector, + ))); + } + Err(e) if connector_tx.borrow().generation() == generation => { + tracing::warn!( + cluster = %cluster_name, + error = %e, + "CDS update rejected; keeping previous connector", + ); + } + Err(e) => { + tracing::warn!( + cluster = %cluster_name, + error = %e, + "re-added CDS update rejected; endpoints remain unavailable", + ); + } } } else => return, @@ -387,6 +426,7 @@ impl Connector for TlsConnector { #[cfg(test)] mod tests { use super::*; + use crate::xds::resource::circuit_breaking::CircuitBreakingConfig; use crate::xds::resource::cluster::{ClusterResource, LbPolicy}; fn plaintext_cluster() -> ClusterResource { @@ -395,6 +435,7 @@ mod tests { eds_service_name: None, lb_policy: LbPolicy::RoundRobin, security: None, + circuit_breaking: CircuitBreakingConfig::default(), } } diff --git a/tonic-xds/src/xds/endpoint_manager.rs b/tonic-xds/src/xds/endpoint_manager.rs index b9153d397..26d228383 100644 --- a/tonic-xds/src/xds/endpoint_manager.rs +++ b/tonic-xds/src/xds/endpoint_manager.rs @@ -33,42 +33,55 @@ use std::collections::HashSet; use std::sync::Arc; -use arc_swap::ArcSwap; -use tokio::sync::mpsc; +use tokio::sync::{mpsc, watch}; use tokio_stream::wrappers::ReceiverStream; use tower::BoxError; use tower::discover::Change; use crate::client::endpoint::{Connector, EndpointAddress}; use crate::client::lb::BoxDiscover; -use crate::xds::cache::CacheWatch; +use crate::xds::cache::{CacheEvent, CacheWatch}; use crate::xds::resource::EndpointsResource; /// Buffer capacity for the endpoint change channel between the diff loop /// and Tower's load balancer. const ENDPOINT_CHANNEL_CAPACITY: usize = 64; -/// An atomically-swappable [`Connector`] held by an [`EndpointManager`]. -/// -/// `XdsClusterDiscovery` stores a snapshot of the cluster's per-CDS-update -/// connector here. The diff loop calls `load_full()` on every new endpoint -/// so each connection picks up the latest snapshot. Existing endpoint -/// channels keep their `EndpointChannel` instance (and any in-flight TLS -/// session) — only freshly-discovered endpoints see the swapped value. -pub(crate) type ConnectorSwap = Arc + Send + Sync>>>; +pub(crate) struct ConnectorState { + generation: u64, + connector: Arc + Send + Sync>, +} + +impl ConnectorState { + pub(crate) fn new( + generation: u64, + connector: Arc + Send + Sync>, + ) -> Self { + Self { + generation, + connector, + } + } + + pub(crate) fn generation(&self) -> u64 { + self.generation + } +} + +pub(crate) type ConnectorWatch = watch::Receiver>>; /// Converts endpoint cache watches into incremental [`Change`] streams. /// /// `EndpointManager` is a pure diff-and-connect component: the caller /// (typically `XdsClusterDiscovery`) obtains a [`CacheWatch`] from the /// [`XdsCache`](crate::xds::cache::XdsCache) and passes it here, plus a -/// [`ConnectorSwap`] that the caller may swap on CDS updates. +/// [`ConnectorWatch`] updated from CDS. pub(crate) struct EndpointManager { - connector: ConnectorSwap, + connector: ConnectorWatch, } impl EndpointManager { - pub(crate) fn new(connector: ConnectorSwap) -> Self { + pub(crate) fn new(connector: ConnectorWatch) -> Self { Self { connector } } @@ -84,9 +97,7 @@ impl EndpointManager { let connector = self.connector.clone(); let (tx, rx) = mpsc::channel(ENDPOINT_CHANNEL_CAPACITY); - // The spawned task exits naturally when either: - // - The CacheWatch closes (cache.remove_endpoints() drops the watch sender) - // - The receiver is dropped (consumer no longer reading Change events) + // The spawned task exits when the cache or the consumer is dropped. tokio::spawn(diff_loop(watch, connector, tx)); Box::pin(ReceiverStream::new(rx)) @@ -100,19 +111,58 @@ impl EndpointManager { /// new endpoints followed by `Remove` for gone ones. async fn diff_loop( mut watch: CacheWatch, - connector: ConnectorSwap, + mut connector: ConnectorWatch, tx: mpsc::Sender, BoxError>>, ) { let mut active: HashSet = HashSet::new(); - while let Some(endpoints) = watch.next().await { + 'events: loop { + let event = tokio::select! { + _ = tx.closed() => return, + event = watch.next_event() => event, + }; + let Some(event) = event else { + return; + }; + let CacheEvent::Resource { + generation, + revision, + resource: endpoints, + } = event + else { + for removed in active.drain() { + if tx.send(Ok(Change::Remove(removed))).await.is_err() { + return; + } + } + continue; + }; + let connector = tokio::select! { + _ = tx.closed() => return, + connector = connector_for_generation(&mut connector, generation) => connector, + }; + let Some(connector) = connector else { + continue; + }; + if watch.current_revision() != revision { + continue; + } + let new_set: HashSet = endpoints .healthy_endpoints() .map(|ep| ep.address.clone()) .collect(); + let added: Vec<_> = new_set.difference(&active).cloned().collect(); + let removed: Vec<_> = active.difference(&new_set).cloned().collect(); - for added in new_set.difference(&active) { - let svc = connector.load_full().connect(added).await; + for added in added { + if watch.current_revision() != revision { + continue 'events; + } + let svc = connector.connect(&added).await; + if watch.current_revision() != revision { + continue 'events; + } if tx .send(Ok(Change::Insert(added.clone(), svc))) .await @@ -120,15 +170,35 @@ async fn diff_loop( { return; } + active.insert(added); } - for removed in active.difference(&new_set) { + for removed in removed { + if watch.current_revision() != revision { + continue 'events; + } if tx.send(Ok(Change::Remove(removed.clone()))).await.is_err() { return; } + active.remove(&removed); } + } +} - active = new_set; +async fn connector_for_generation( + connector: &mut ConnectorWatch, + generation: u64, +) -> Option + Send + Sync>> { + loop { + let current = connector.borrow().clone(); + match current.generation.cmp(&generation) { + std::cmp::Ordering::Equal => return Some(Arc::clone(¤t.connector)), + std::cmp::Ordering::Greater => return None, + std::cmp::Ordering::Less => {} + } + if connector.changed().await.is_err() { + return None; + } } } @@ -152,9 +222,18 @@ mod tests { } } - fn test_swap() -> ConnectorSwap { + fn test_connector_channel( + generation: u64, + ) -> ( + watch::Sender>>, + ConnectorWatch, + ) { let conn: Arc + Send + Sync> = Arc::new(StringConnector); - Arc::new(ArcSwap::from_pointee(conn)) + watch::channel(Arc::new(ConnectorState::new(generation, conn))) + } + + fn test_connector_watch() -> ConnectorWatch { + test_connector_channel(0).1 } fn make_endpoints(cluster: &str, addrs: &[(&str, u16)]) -> Arc { @@ -179,7 +258,7 @@ mod tests { #[tokio::test] async fn initial_endpoints_emitted_as_inserts() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let manager = EndpointManager::new(test_connector_watch()); cache.update_endpoints( "c1", @@ -202,7 +281,7 @@ mod tests { #[tokio::test] async fn added_endpoint_emits_insert() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let manager = EndpointManager::new(test_connector_watch()); cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); @@ -223,7 +302,7 @@ mod tests { #[tokio::test] async fn removed_endpoint_emits_remove() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let manager = EndpointManager::new(test_connector_watch()); cache.update_endpoints( "c1", @@ -247,7 +326,7 @@ mod tests { #[tokio::test] async fn unhealthy_endpoint_removed() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let manager = EndpointManager::new(test_connector_watch()); cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); @@ -276,9 +355,10 @@ mod tests { } #[tokio::test] - async fn cache_removal_closes_stream() { + async fn cache_removal_clears_endpoints_and_readd_recovers() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let (connector_tx, connector_rx) = test_connector_channel(0); + let manager = EndpointManager::new(connector_rx); cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); @@ -287,13 +367,58 @@ mod tests { cache.remove_endpoints("c1"); - assert!(stream.next().await.is_none()); + match stream.next().await.unwrap().unwrap() { + Change::Remove(addr) => assert_eq!(addr.to_string(), "10.0.0.1:8080"), + Change::Insert(..) => panic!("expected Remove after cache removal"), + } + + let connector: Arc + Send + Sync> = + Arc::new(StringConnector); + connector_tx.send_replace(Arc::new(ConnectorState::new(1, connector))); + cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); + match stream.next().await.unwrap().unwrap() { + Change::Insert(addr, _) => assert_eq!(addr.to_string(), "10.0.0.1:8080"), + Change::Remove(..) => panic!("expected Insert after cache re-add"), + } + } + + #[tokio::test] + async fn readded_endpoints_wait_for_matching_connector_generation() { + let cache = XdsCache::new(); + let (connector_tx, connector_rx) = test_connector_channel(0); + let manager = EndpointManager::new(connector_rx); + cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); + let mut stream = manager.discover_endpoints(cache.watch_endpoints("c1")); + let _ = stream.next().await; // consume initial insert + + cache.remove_endpoints("c1"); + cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.2", 8080)])); + + match stream.next().await.unwrap().unwrap() { + Change::Remove(addr) => assert_eq!(addr.to_string(), "10.0.0.1:8080"), + Change::Insert(..) => panic!("expected old endpoint removal"), + } + assert!( + tokio::time::timeout(std::time::Duration::from_millis(20), stream.next()) + .await + .is_err(), + "re-added endpoints must wait for their CDS connector generation", + ); + + cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.3", 8080)])); + let connector: Arc + Send + Sync> = + Arc::new(StringConnector); + connector_tx.send_replace(Arc::new(ConnectorState::new(1, connector))); + match stream.next().await.unwrap().unwrap() { + Change::Insert(addr, _) => assert_eq!(addr.to_string(), "10.0.0.3:8080"), + Change::Remove(..) => panic!("expected re-added endpoint insert"), + } } #[tokio::test] async fn multiple_clusters_independent() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let manager = EndpointManager::new(test_connector_watch()); cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); cache.update_endpoints("c2", make_endpoints("c2", &[("10.0.0.2", 9090)])); @@ -314,7 +439,7 @@ mod tests { #[tokio::test] async fn endpoint_swap_emits_insert_then_remove() { let cache = XdsCache::new(); - let manager = EndpointManager::new(test_swap()); + let manager = EndpointManager::new(test_connector_watch()); cache.update_endpoints("c1", make_endpoints("c1", &[("10.0.0.1", 8080)])); diff --git a/tonic-xds/src/xds/resource/circuit_breaking.rs b/tonic-xds/src/xds/resource/circuit_breaking.rs index 916aed927..c8c69f48b 100644 --- a/tonic-xds/src/xds/resource/circuit_breaking.rs +++ b/tonic-xds/src/xds/resource/circuit_breaking.rs @@ -29,9 +29,7 @@ //! because they are connection-pool or retry specific and do not apply to gRPC's //! A32 request limiter. //! -//! Client-side limiter primitives can consume this config; production CDS wiring -//! is added separately so validation only advertises support once enforcement is -//! in the request path. +//! `ClusterResource` carries this config to the client-side limiter. //! //! [gRFC A32]: https://github.com/grpc/proposal/blob/master/A32-xds-circuit-breaking.md diff --git a/tonic-xds/src/xds/resource/cluster.rs b/tonic-xds/src/xds/resource/cluster.rs index 618f8440c..848c1aa2a 100644 --- a/tonic-xds/src/xds/resource/cluster.rs +++ b/tonic-xds/src/xds/resource/cluster.rs @@ -30,6 +30,7 @@ use prost::Message; use xds_client::resource::TypeUrl; use xds_client::{Error, Resource}; +use super::circuit_breaking::CircuitBreakingConfig; use super::security::{ClusterSecurityConfig, parse_transport_socket}; /// Validated Cluster resource. @@ -44,6 +45,8 @@ pub(crate) struct ClusterResource { /// TLS security config parsed from `transport_socket`. `None` means the /// cluster uses plaintext connections. pub security: Option, + /// Circuit-breaking config for the cluster. + pub circuit_breaking: CircuitBreakingConfig, } /// Load balancing policies. @@ -91,12 +94,14 @@ impl Resource for ClusterResource { }; let security = parse_transport_socket(message.transport_socket)?; + let circuit_breaking = CircuitBreakingConfig::from_proto(message.circuit_breakers.as_ref()); Ok(ClusterResource { name, eds_service_name, lb_policy, security, + circuit_breaking, }) } } @@ -129,6 +134,7 @@ mod tests { assert_eq!(validated.name, "my-cluster"); assert_eq!(validated.lb_policy, LbPolicy::RoundRobin); assert!(validated.eds_service_name.is_none()); + assert_eq!(validated.circuit_breaking, CircuitBreakingConfig::default()); } #[test] @@ -165,6 +171,34 @@ mod tests { assert_eq!(validated.lb_policy, LbPolicy::LeastRequest); } + #[test] + fn test_circuit_breaking_config() { + use envoy_types::pb::envoy::config::cluster::v3::CircuitBreakers; + use envoy_types::pb::envoy::config::cluster::v3::circuit_breakers::Thresholds; + use envoy_types::pb::envoy::config::core::v3::RoutingPriority; + use envoy_types::pb::google::protobuf::UInt32Value; + + let cluster = Cluster { + name: "cb-cluster".to_string(), + lb_policy: cluster::LbPolicy::RoundRobin as i32, + circuit_breakers: Some(CircuitBreakers { + thresholds: vec![Thresholds { + priority: RoutingPriority::Default as i32, + max_requests: Some(UInt32Value { value: 7 }), + ..Default::default() + }], + ..Default::default() + }), + ..Default::default() + }; + + let validated = ClusterResource::validate(cluster).unwrap(); + assert_eq!( + validated.circuit_breaking, + CircuitBreakingConfig { max_requests: 7 }, + ); + } + #[test] fn test_unsupported_lb_policy_is_rejected() { let cluster = Cluster {