1use async_trait::async_trait;
3use ecat_circuit_breaker::{Breaker, BreakerConfig, BreakerError, BreakerState};
4use ecat_errors::{Error, ErrorCode};
5
6use crate::dialect::Dialect;
7use crate::rdbms::{RdbmsClient, RdbmsError, Row, SqlExecutor, Transaction};
8
9pub struct CircuitBreakerExecutor<S> {
14 inner: S,
15 breaker: Breaker,
16}
17
18impl<S> CircuitBreakerExecutor<S> {
19 pub fn new(inner: S, cfg: BreakerConfig) -> Self {
20 Self {
21 inner,
22 breaker: Breaker::new(cfg),
23 }
24 }
25
26 pub fn state(&self) -> BreakerState {
28 self.breaker.state()
29 }
30
31 pub fn inner(&self) -> &S {
32 &self.inner
33 }
34}
35
36pub fn map_breaker_error(e: BreakerError<RdbmsError>) -> RdbmsError {
40 match e {
41 BreakerError::Inner(inner) => inner,
42 other => RdbmsError::Connection(other.to_string()),
43 }
44}
45
46pub fn breaker_error_to_backend_error(e: BreakerError<Error>, reason: &'static str) -> Error {
52 match e {
53 BreakerError::Inner(inner) => inner,
54 other => Error::new(ErrorCode::Unavailable, reason, other.to_string()),
55 }
56}
57
58#[async_trait]
59impl<S: SqlExecutor> SqlExecutor for CircuitBreakerExecutor<S> {
60 async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
61 self.breaker
63 .call(|| self.inner.execute(sql))
64 .await
65 .map_err(map_breaker_error)
66 }
67
68 async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
69 self.breaker
70 .call(|| self.inner.query(sql))
71 .await
72 .map_err(map_breaker_error)
73 }
74
75 async fn execute_with(
76 &self,
77 sql: &str,
78 params: &[serde_json::Value],
79 ) -> Result<u64, RdbmsError> {
80 self.breaker
81 .call(|| self.inner.execute_with(sql, params))
82 .await
83 .map_err(map_breaker_error)
84 }
85
86 async fn query_with(
87 &self,
88 sql: &str,
89 params: &[serde_json::Value],
90 ) -> Result<Vec<Row>, RdbmsError> {
91 self.breaker
92 .call(|| self.inner.query_with(sql, params))
93 .await
94 .map_err(map_breaker_error)
95 }
96
97 async fn query_write(
99 &self,
100 sql: &str,
101 params: &[serde_json::Value],
102 ) -> Result<Vec<Row>, RdbmsError> {
103 self.breaker
104 .call(|| self.inner.query_write(sql, params))
105 .await
106 .map_err(map_breaker_error)
107 }
108
109 async fn execute_then_query(
110 &self,
111 first: &str,
112 first_params: &[serde_json::Value],
113 second: &str,
114 ) -> Result<Vec<Row>, RdbmsError> {
115 self.breaker
116 .call(|| self.inner.execute_then_query(first, first_params, second))
117 .await
118 .map_err(map_breaker_error)
119 }
120
121 fn dialect(&self) -> Dialect {
123 self.inner.dialect()
124 }
125}
126
127#[async_trait]
128impl<S: RdbmsClient> RdbmsClient for CircuitBreakerExecutor<S> {
129 async fn transaction(&self) -> Result<Transaction, RdbmsError> {
136 self.inner.transaction().await
137 }
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143 use crate::dialect::Dialect;
144 use crate::rdbms::{RdbmsClient, RdbmsError, Row, SqlExecutor, Transaction};
145 use async_trait::async_trait;
146 use ecat_circuit_breaker::{BreakerConfig, BreakerError, BreakerState};
147 use ecat_errors::{Error, ErrorCode};
148 use std::sync::Arc;
149 use std::sync::atomic::{AtomicUsize, Ordering};
150
151 struct FakeExecutor {
154 calls: Arc<AtomicUsize>,
155 fail: bool,
156 }
157
158 impl FakeExecutor {
159 fn new(fail: bool) -> (Self, Arc<AtomicUsize>) {
160 let calls = Arc::new(AtomicUsize::new(0));
161 (
162 Self {
163 calls: Arc::clone(&calls),
164 fail,
165 },
166 calls,
167 )
168 }
169
170 fn record<T>(&self, ok: T) -> Result<T, RdbmsError> {
171 self.calls.fetch_add(1, Ordering::SeqCst);
172 if self.fail {
173 Err(RdbmsError::Connection("backend down".into()))
174 } else {
175 Ok(ok)
176 }
177 }
178 }
179
180 #[async_trait]
181 impl SqlExecutor for FakeExecutor {
182 async fn execute(&self, _sql: &str) -> Result<u64, RdbmsError> {
183 self.record(1)
184 }
185 async fn query(&self, _sql: &str) -> Result<Vec<Row>, RdbmsError> {
186 self.record(vec![])
187 }
188 async fn execute_with(
189 &self,
190 _sql: &str,
191 _params: &[serde_json::Value],
192 ) -> Result<u64, RdbmsError> {
193 self.record(1)
194 }
195 async fn query_with(
196 &self,
197 _sql: &str,
198 _params: &[serde_json::Value],
199 ) -> Result<Vec<Row>, RdbmsError> {
200 self.record(vec![])
201 }
202 async fn query_write(
203 &self,
204 _sql: &str,
205 _params: &[serde_json::Value],
206 ) -> Result<Vec<Row>, RdbmsError> {
207 self.record(vec![])
208 }
209 async fn execute_then_query(
210 &self,
211 _first: &str,
212 _first_params: &[serde_json::Value],
213 _second: &str,
214 ) -> Result<Vec<Row>, RdbmsError> {
215 self.record(vec![])
216 }
217 fn dialect(&self) -> Dialect {
218 Dialect::Sqlite
219 }
220 }
221
222 #[async_trait]
223 impl RdbmsClient for FakeExecutor {
224 async fn transaction(&self) -> Result<Transaction, RdbmsError> {
225 self.record(Transaction::new())
226 }
227 }
228
229 #[tokio::test]
234 async fn transaction_delegates_without_going_through_the_breaker() {
235 let (inner, calls) = FakeExecutor::new(true);
236 let exec = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
237
238 assert!(exec.transaction().await.is_err());
239 assert_eq!(calls.load(Ordering::SeqCst), 1, "必须委托给内层");
240 assert_eq!(exec.state(), BreakerState::Closed);
241
242 for _ in 0..5 {
244 assert!(exec.transaction().await.is_err());
245 }
246 assert_eq!(calls.load(Ordering::SeqCst), 6);
247 assert_eq!(
248 exec.state(),
249 BreakerState::Closed,
250 "事务失败不得计入熔断窗口(熔断器根本没被碰)"
251 );
252 }
253
254 #[tokio::test]
257 async fn open_breaker_never_reaches_the_inner_executor() {
258 let (inner, calls) = FakeExecutor::new(true);
259 let exec = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
260
261 for _ in 0..5 {
263 assert!(exec.query("SELECT 1").await.is_err());
264 }
265 assert_eq!(exec.state(), BreakerState::Open);
266
267 calls.store(0, Ordering::SeqCst);
269 assert!(exec.execute("UPDATE t SET x = 1").await.is_err());
270 assert!(exec.query("SELECT 1").await.is_err());
271 assert!(exec.execute_with("UPDATE t SET x = ?", &[]).await.is_err());
272 assert!(exec.query_with("SELECT ?", &[]).await.is_err());
273 assert!(
274 exec.query_write("INSERT INTO t (v) VALUES (?) RETURNING id", &[])
275 .await
276 .is_err()
277 );
278 let err = exec
279 .execute_then_query(
280 "INSERT INTO t (v) VALUES (?)",
281 &[],
282 "SELECT LAST_INSERT_ID()",
283 )
284 .await
285 .unwrap_err();
286 assert!(err.to_string().contains("circuit breaker"), "got: {err}");
288
289 assert_eq!(
290 calls.load(Ordering::SeqCst),
291 0,
292 "熔断打开后内层不得被调用(否则只是快速失败的转发)"
293 );
294 }
295
296 #[tokio::test]
298 async fn inner_errors_count_as_failures() {
299 let (inner, calls) = FakeExecutor::new(false);
300 let ok = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
301 for _ in 0..6 {
302 ok.query_write("INSERT INTO t (v) VALUES (?) RETURNING id", &[])
303 .await
304 .unwrap();
305 }
306 assert_eq!(ok.state(), BreakerState::Closed, "成功不得触发熔断");
307 assert_eq!(calls.load(Ordering::SeqCst), 6);
308
309 let (inner, _) = FakeExecutor::new(true);
310 let bad = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
311 for _ in 0..5 {
312 assert!(bad.execute_with("UPDATE t SET x = ?", &[]).await.is_err());
313 }
314 assert_eq!(bad.state(), BreakerState::Open, "后端报错必须计入失败");
315 }
316
317 #[tokio::test]
320 async fn breaker_state_is_readable() {
321 let (inner, _) = FakeExecutor::new(true);
322 let exec = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
323 assert_eq!(exec.state(), BreakerState::Closed);
324
325 for _ in 0..5 {
326 assert!(exec.execute("UPDATE t SET x = 1").await.is_err());
327 }
328 assert_eq!(exec.state(), BreakerState::Open);
329 assert_eq!(exec.dialect(), Dialect::Sqlite);
330 assert_eq!(exec.inner().dialect(), Dialect::Sqlite);
331 }
332
333 #[test]
338 fn breaker_error_to_backend_error_maps_all_three_variants() {
339 let inner = Error::new(ErrorCode::DeadlineExceeded, "redis", "connect timed out");
340 let passthrough = breaker_error_to_backend_error(BreakerError::Inner(inner), "redis");
341 assert_eq!(passthrough.code, ErrorCode::DeadlineExceeded);
342 assert_eq!(passthrough.reason, "redis");
343 assert_eq!(passthrough.message, "connect timed out");
344
345 let open = breaker_error_to_backend_error(BreakerError::Open, "clickhouse");
346 assert_eq!(open.code, ErrorCode::Unavailable);
347 assert_eq!(open.reason, "clickhouse");
348 assert_eq!(open.message, "circuit breaker is open");
349
350 let probes = breaker_error_to_backend_error(BreakerError::ProbesExhausted, "clickhouse");
351 assert_eq!(probes.code, ErrorCode::Unavailable);
352 assert_eq!(probes.reason, "clickhouse");
353 assert_eq!(probes.message, "circuit breaker: too many probes");
354 }
355}