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