Skip to main content

rpc/
resilience.rs

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
9/// Returns whether a gRPC result should be treated as healthy by a circuit breaker.
10///
11/// Client mistakes and domain failures do not indicate an unhealthy dependency. Transport and
12/// server failures do, matching go-zero's zrpc outcome classification.
13pub 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
25/// Maps a completed gRPC call to breaker health without treating caller cancellation as either a
26/// dependency success or failure.
27pub 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/// Protocol-aware circuit breaking for unary Tonic client calls.
38#[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    /// Runs a unary call when the circuit permits it.
59    ///
60    /// Infrastructure failures count against the circuit. Statuses such as `InvalidArgument`,
61    /// `NotFound`, and `PermissionDenied` are returned without marking the dependency unhealthy.
62    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/// Adaptive admission control for unary Tonic server handlers.
86#[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    /// Admits and observes a unary handler, or rejects it with `ResourceExhausted`.
107    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}