1use std::{future::Future, sync::Arc};
2
3use rust_zero_core::{
4 AdaptiveShedder, BreakerState, CircuitBreaker, CircuitBreakerConfig, CircuitBreakerError,
5 CircuitBreakerSnapshot, CircuitOutcome, LoadShedderConfig,
6};
7use tonic::{Code, Status};
8
9pub fn acceptable_status(status: &Status) -> bool {
14 !matches!(
15 status.code(),
16 Code::DeadlineExceeded
17 | Code::Internal
18 | Code::Unavailable
19 | Code::DataLoss
20 | Code::Unimplemented
21 | Code::ResourceExhausted
22 )
23}
24
25pub fn circuit_outcome(status: &Status) -> CircuitOutcome {
28 if status.code() == Code::Cancelled {
29 CircuitOutcome::Cancellation
30 } else if acceptable_status(status) {
31 CircuitOutcome::Success
32 } else {
33 CircuitOutcome::Failure
34 }
35}
36
37#[derive(Clone)]
39pub struct RpcCircuitBreaker {
40 breaker: Arc<CircuitBreaker>,
41}
42
43impl RpcCircuitBreaker {
44 pub fn new(config: CircuitBreakerConfig) -> Self {
45 Self {
46 breaker: Arc::new(CircuitBreaker::new(config)),
47 }
48 }
49
50 pub fn state(&self) -> BreakerState {
51 self.breaker.state()
52 }
53
54 pub fn snapshot(&self) -> CircuitBreakerSnapshot {
55 self.breaker.snapshot()
56 }
57
58 pub async fn call<T, F, Fut>(&self, operation: F) -> Result<T, Status>
63 where
64 F: FnOnce() -> Fut,
65 Fut: Future<Output = Result<T, Status>>,
66 {
67 match self
68 .breaker
69 .execute_async_with_outcome(operation, |result| {
70 result
71 .as_ref()
72 .map_or_else(circuit_outcome, |_| CircuitOutcome::Success)
73 })
74 .await
75 {
76 Ok(value) => Ok(value),
77 Err(CircuitBreakerError::Operation(status)) => Err(status),
78 Err(CircuitBreakerError::Open) => {
79 Err(Status::unavailable("gRPC dependency circuit is open"))
80 }
81 }
82 }
83}
84
85#[derive(Clone)]
87pub struct RpcLoadShedder {
88 shedder: AdaptiveShedder,
89}
90
91impl RpcLoadShedder {
92 pub fn new(config: LoadShedderConfig) -> Self {
93 Self {
94 shedder: AdaptiveShedder::new(config),
95 }
96 }
97
98 pub fn current_limit(&self) -> usize {
99 self.shedder.current_limit()
100 }
101
102 pub fn in_flight(&self) -> usize {
103 self.shedder.in_flight()
104 }
105
106 pub async fn call<T, F, Fut>(&self, operation: F) -> Result<T, Status>
108 where
109 F: FnOnce() -> Fut,
110 Fut: Future<Output = Result<T, Status>>,
111 {
112 let _permit = self
113 .shedder
114 .try_acquire()
115 .ok_or_else(|| Status::resource_exhausted("gRPC server is overloaded"))?;
116 operation().await
117 }
118}
119
120#[cfg(test)]
121mod tests {
122 use super::*;
123 use std::{
124 sync::atomic::{AtomicBool, Ordering},
125 time::Duration,
126 };
127 use tokio::sync::Notify;
128
129 #[test]
130 fn classifies_application_and_infrastructure_statuses() {
131 for code in [
132 Code::DeadlineExceeded,
133 Code::Internal,
134 Code::Unavailable,
135 Code::DataLoss,
136 Code::Unimplemented,
137 Code::ResourceExhausted,
138 ] {
139 assert!(!acceptable_status(&Status::new(code, "failure")));
140 }
141
142 for code in [
143 Code::InvalidArgument,
144 Code::NotFound,
145 Code::AlreadyExists,
146 Code::PermissionDenied,
147 Code::Unauthenticated,
148 ] {
149 assert!(acceptable_status(&Status::new(code, "application error")));
150 }
151
152 assert_eq!(
153 circuit_outcome(&Status::cancelled("caller left")),
154 CircuitOutcome::Cancellation
155 );
156 }
157
158 #[tokio::test]
159 async fn circuit_opens_only_for_infrastructure_failures() {
160 let breaker = RpcCircuitBreaker::new(CircuitBreakerConfig::new(2, Duration::from_secs(30)));
161
162 for _ in 0..3 {
163 let error = breaker
164 .call(|| async { Err::<(), _>(Status::invalid_argument("bad request")) })
165 .await
166 .unwrap_err();
167 assert_eq!(error.code(), Code::InvalidArgument);
168 }
169 assert_eq!(breaker.state(), BreakerState::Closed);
170
171 for _ in 0..2 {
172 let error = breaker
173 .call(|| async { Err::<(), _>(Status::unavailable("offline")) })
174 .await
175 .unwrap_err();
176 assert_eq!(error.message(), "offline");
177 }
178 assert_eq!(breaker.state(), BreakerState::Open);
179
180 let invoked = AtomicBool::new(false);
181 let error = breaker
182 .call(|| async {
183 invoked.store(true, Ordering::Relaxed);
184 Ok(())
185 })
186 .await
187 .unwrap_err();
188 assert_eq!(error.code(), Code::Unavailable);
189 assert!(!invoked.load(Ordering::Relaxed));
190 }
191
192 #[tokio::test]
193 async fn rolling_breaker_tracks_protocol_outcomes_and_rejects_faulty_traffic() {
194 use rust_zero_core::RollingCircuitBreakerConfig;
195
196 let breaker = RpcCircuitBreaker::new(CircuitBreakerConfig::rolling(
197 RollingCircuitBreakerConfig::new()
198 .with_minimum_requests(1)
199 .with_sensitivity(0.1)
200 .with_random_seed(3),
201 ));
202 let cancelled = breaker
203 .call(|| async { Err::<(), _>(Status::cancelled("caller left")) })
204 .await
205 .unwrap_err();
206 assert_eq!(cancelled.code(), Code::Cancelled);
207
208 for _ in 0..2 {
209 let _ = breaker
210 .call(|| async { Err::<(), _>(Status::unavailable("offline")) })
211 .await;
212 }
213 let snapshot = breaker.snapshot();
214 assert_eq!(snapshot.cancellations, 1);
215 assert_eq!(snapshot.failures, 2);
216 assert_eq!(snapshot.total, 2);
217 assert!(snapshot.drop_ratio > 0.0);
218
219 for _ in 0..20 {
220 let _ = breaker.call(|| async { Ok::<_, Status>(()) }).await;
221 }
222 assert!(breaker.snapshot().rejections > 0);
223 }
224
225 #[tokio::test]
226 async fn load_shedder_rejects_work_beyond_the_current_limit() {
227 let shedder = RpcLoadShedder::new(LoadShedderConfig::new(1, Duration::from_secs(1)));
228 let entered = Arc::new(Notify::new());
229 let release = Arc::new(Notify::new());
230 let task_shedder = shedder.clone();
231 let task_entered = Arc::clone(&entered);
232 let task_release = Arc::clone(&release);
233 let active = tokio::spawn(async move {
234 task_shedder
235 .call(|| async move {
236 task_entered.notify_one();
237 task_release.notified().await;
238 Ok(())
239 })
240 .await
241 });
242
243 entered.notified().await;
244 assert_eq!(shedder.in_flight(), 1);
245 let error = shedder.call(|| async { Ok(()) }).await.unwrap_err();
246 assert_eq!(error.code(), Code::ResourceExhausted);
247
248 release.notify_one();
249 active.await.unwrap().unwrap();
250 assert_eq!(shedder.in_flight(), 0);
251 }
252}