use async_trait::async_trait;
use ecat_circuit_breaker::{Breaker, BreakerConfig, BreakerError, BreakerState};
use ecat_errors::{Error, ErrorCode};
use crate::dialect::Dialect;
use crate::rdbms::{RdbmsClient, RdbmsError, Row, SqlExecutor, Transaction};
pub struct CircuitBreakerExecutor<S> {
inner: S,
breaker: Breaker,
}
impl<S> CircuitBreakerExecutor<S> {
pub fn new(inner: S, cfg: BreakerConfig) -> Self {
Self {
inner,
breaker: Breaker::new(cfg),
}
}
pub fn state(&self) -> BreakerState {
self.breaker.state()
}
pub fn inner(&self) -> &S {
&self.inner
}
}
pub fn map_breaker_error(e: BreakerError<RdbmsError>) -> RdbmsError {
match e {
BreakerError::Inner(inner) => inner,
other => RdbmsError::Connection(other.to_string()),
}
}
pub fn breaker_error_to_backend_error(e: BreakerError<Error>, reason: &'static str) -> Error {
match e {
BreakerError::Inner(inner) => inner,
other => Error::new(ErrorCode::Unavailable, reason, other.to_string()),
}
}
#[async_trait]
impl<S: SqlExecutor> SqlExecutor for CircuitBreakerExecutor<S> {
async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
self.breaker
.call(|| self.inner.execute(sql))
.await
.map_err(map_breaker_error)
}
async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
self.breaker
.call(|| self.inner.query(sql))
.await
.map_err(map_breaker_error)
}
async fn execute_with(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<u64, RdbmsError> {
self.breaker
.call(|| self.inner.execute_with(sql, params))
.await
.map_err(map_breaker_error)
}
async fn query_with(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
self.breaker
.call(|| self.inner.query_with(sql, params))
.await
.map_err(map_breaker_error)
}
async fn query_write(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
self.breaker
.call(|| self.inner.query_write(sql, params))
.await
.map_err(map_breaker_error)
}
async fn execute_then_query(
&self,
first: &str,
first_params: &[serde_json::Value],
second: &str,
) -> Result<Vec<Row>, RdbmsError> {
self.breaker
.call(|| self.inner.execute_then_query(first, first_params, second))
.await
.map_err(map_breaker_error)
}
fn dialect(&self) -> Dialect {
self.inner.dialect()
}
}
#[async_trait]
impl<S: RdbmsClient> RdbmsClient for CircuitBreakerExecutor<S> {
async fn transaction(&self) -> Result<Transaction, RdbmsError> {
self.inner.transaction().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dialect::Dialect;
use crate::rdbms::{RdbmsClient, RdbmsError, Row, SqlExecutor, Transaction};
use async_trait::async_trait;
use ecat_circuit_breaker::{BreakerConfig, BreakerError, BreakerState};
use ecat_errors::{Error, ErrorCode};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
struct FakeExecutor {
calls: Arc<AtomicUsize>,
fail: bool,
}
impl FakeExecutor {
fn new(fail: bool) -> (Self, Arc<AtomicUsize>) {
let calls = Arc::new(AtomicUsize::new(0));
(
Self {
calls: Arc::clone(&calls),
fail,
},
calls,
)
}
fn record<T>(&self, ok: T) -> Result<T, RdbmsError> {
self.calls.fetch_add(1, Ordering::SeqCst);
if self.fail {
Err(RdbmsError::Connection("backend down".into()))
} else {
Ok(ok)
}
}
}
#[async_trait]
impl SqlExecutor for FakeExecutor {
async fn execute(&self, _sql: &str) -> Result<u64, RdbmsError> {
self.record(1)
}
async fn query(&self, _sql: &str) -> Result<Vec<Row>, RdbmsError> {
self.record(vec![])
}
async fn execute_with(
&self,
_sql: &str,
_params: &[serde_json::Value],
) -> Result<u64, RdbmsError> {
self.record(1)
}
async fn query_with(
&self,
_sql: &str,
_params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
self.record(vec![])
}
async fn query_write(
&self,
_sql: &str,
_params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
self.record(vec![])
}
async fn execute_then_query(
&self,
_first: &str,
_first_params: &[serde_json::Value],
_second: &str,
) -> Result<Vec<Row>, RdbmsError> {
self.record(vec![])
}
fn dialect(&self) -> Dialect {
Dialect::Sqlite
}
}
#[async_trait]
impl RdbmsClient for FakeExecutor {
async fn transaction(&self) -> Result<Transaction, RdbmsError> {
self.record(Transaction::new())
}
}
#[tokio::test]
async fn transaction_delegates_without_going_through_the_breaker() {
let (inner, calls) = FakeExecutor::new(true);
let exec = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
assert!(exec.transaction().await.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1, "必须委托给内层");
assert_eq!(exec.state(), BreakerState::Closed);
for _ in 0..5 {
assert!(exec.transaction().await.is_err());
}
assert_eq!(calls.load(Ordering::SeqCst), 6);
assert_eq!(
exec.state(),
BreakerState::Closed,
"事务失败不得计入熔断窗口(熔断器根本没被碰)"
);
}
#[tokio::test]
async fn open_breaker_never_reaches_the_inner_executor() {
let (inner, calls) = FakeExecutor::new(true);
let exec = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
for _ in 0..5 {
assert!(exec.query("SELECT 1").await.is_err());
}
assert_eq!(exec.state(), BreakerState::Open);
calls.store(0, Ordering::SeqCst);
assert!(exec.execute("UPDATE t SET x = 1").await.is_err());
assert!(exec.query("SELECT 1").await.is_err());
assert!(exec.execute_with("UPDATE t SET x = ?", &[]).await.is_err());
assert!(exec.query_with("SELECT ?", &[]).await.is_err());
assert!(
exec.query_write("INSERT INTO t (v) VALUES (?) RETURNING id", &[])
.await
.is_err()
);
let err = exec
.execute_then_query(
"INSERT INTO t (v) VALUES (?)",
&[],
"SELECT LAST_INSERT_ID()",
)
.await
.unwrap_err();
assert!(err.to_string().contains("circuit breaker"), "got: {err}");
assert_eq!(
calls.load(Ordering::SeqCst),
0,
"熔断打开后内层不得被调用(否则只是快速失败的转发)"
);
}
#[tokio::test]
async fn inner_errors_count_as_failures() {
let (inner, calls) = FakeExecutor::new(false);
let ok = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
for _ in 0..6 {
ok.query_write("INSERT INTO t (v) VALUES (?) RETURNING id", &[])
.await
.unwrap();
}
assert_eq!(ok.state(), BreakerState::Closed, "成功不得触发熔断");
assert_eq!(calls.load(Ordering::SeqCst), 6);
let (inner, _) = FakeExecutor::new(true);
let bad = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
for _ in 0..5 {
assert!(bad.execute_with("UPDATE t SET x = ?", &[]).await.is_err());
}
assert_eq!(bad.state(), BreakerState::Open, "后端报错必须计入失败");
}
#[tokio::test]
async fn breaker_state_is_readable() {
let (inner, _) = FakeExecutor::new(true);
let exec = CircuitBreakerExecutor::new(inner, BreakerConfig::default());
assert_eq!(exec.state(), BreakerState::Closed);
for _ in 0..5 {
assert!(exec.execute("UPDATE t SET x = 1").await.is_err());
}
assert_eq!(exec.state(), BreakerState::Open);
assert_eq!(exec.dialect(), Dialect::Sqlite);
assert_eq!(exec.inner().dialect(), Dialect::Sqlite);
}
#[test]
fn breaker_error_to_backend_error_maps_all_three_variants() {
let inner = Error::new(ErrorCode::DeadlineExceeded, "redis", "connect timed out");
let passthrough = breaker_error_to_backend_error(BreakerError::Inner(inner), "redis");
assert_eq!(passthrough.code, ErrorCode::DeadlineExceeded);
assert_eq!(passthrough.reason, "redis");
assert_eq!(passthrough.message, "connect timed out");
let open = breaker_error_to_backend_error(BreakerError::Open, "clickhouse");
assert_eq!(open.code, ErrorCode::Unavailable);
assert_eq!(open.reason, "clickhouse");
assert_eq!(open.message, "circuit breaker is open");
let probes = breaker_error_to_backend_error(BreakerError::ProbesExhausted, "clickhouse");
assert_eq!(probes.code, ErrorCode::Unavailable);
assert_eq!(probes.reason, "clickhouse");
assert_eq!(probes.message, "circuit breaker: too many probes");
}
}