use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use ecat_circuit_breaker::{Breaker, BreakerConfig, BreakerState};
use crate::breaker::map_breaker_error;
use crate::dialect::Dialect;
use crate::rdbms::{RdbmsClient, RdbmsError, Row, SqlExecutor, Transaction};
struct Endpoint {
client: Arc<dyn RdbmsClient>,
breaker: Breaker,
}
impl Endpoint {
fn new(client: Arc<dyn RdbmsClient>, cfg: &BreakerConfig) -> Self {
Self {
client,
breaker: Breaker::new(cfg.clone()),
}
}
fn is_available(&self) -> bool {
self.breaker.state() != BreakerState::Open
}
fn dialect(&self) -> Dialect {
self.client.dialect()
}
async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
self.breaker
.call(|| self.client.execute(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.client.execute_with(sql, params))
.await
.map_err(map_breaker_error)
}
async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
self.breaker
.call(|| self.client.query(sql))
.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.client.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.client.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.client.execute_then_query(first, first_params, second))
.await
.map_err(map_breaker_error)
}
}
pub struct RdbmsRouting {
primary: Endpoint,
replicas: Vec<Endpoint>,
next: AtomicUsize,
fallback_to_primary: bool,
}
impl RdbmsRouting {
pub fn new(primary: Arc<dyn RdbmsClient>, replicas: Vec<Arc<dyn RdbmsClient>>) -> Self {
Self::with_breaker_config(primary, replicas, BreakerConfig::default())
}
pub fn with_breaker_config(
primary: Arc<dyn RdbmsClient>,
replicas: Vec<Arc<dyn RdbmsClient>>,
cfg: BreakerConfig,
) -> Self {
Self {
primary: Endpoint::new(primary, &cfg),
replicas: replicas
.into_iter()
.map(|client| Endpoint::new(client, &cfg))
.collect(),
next: AtomicUsize::new(0),
fallback_to_primary: true,
}
}
pub fn fallback_to_primary(mut self, yes: bool) -> Self {
self.fallback_to_primary = yes;
self
}
fn pick_replica(&self) -> Option<&Endpoint> {
let n = self.replicas.len();
if n == 0 {
return None;
}
let start = self.next.fetch_add(1, Ordering::Relaxed);
(0..n)
.map(|i| &self.replicas[(start + i) % n])
.find(|ep| ep.is_available())
}
}
#[async_trait]
impl SqlExecutor for RdbmsRouting {
async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
self.primary.execute(sql).await
}
async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
match self.pick_replica() {
Some(replica) => replica.query(sql).await,
None if self.fallback_to_primary => self.primary.query(sql).await,
None => Err(RdbmsError::NoAvailableReplica),
}
}
async fn execute_with(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<u64, RdbmsError> {
self.primary.execute_with(sql, params).await
}
async fn query_with(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
match self.pick_replica() {
Some(replica) => replica.query_with(sql, params).await,
None if self.fallback_to_primary => self.primary.query_with(sql, params).await,
None => Err(RdbmsError::NoAvailableReplica),
}
}
async fn query_write(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
self.primary.query_write(sql, params).await
}
async fn execute_then_query(
&self,
first: &str,
first_params: &[serde_json::Value],
second: &str,
) -> Result<Vec<Row>, RdbmsError> {
self.primary
.execute_then_query(first, first_params, second)
.await
}
fn dialect(&self) -> Dialect {
self.primary.dialect()
}
}
#[async_trait]
impl RdbmsClient for RdbmsRouting {
async fn transaction(&self) -> Result<Transaction, RdbmsError> {
self.primary.client.transaction().await
}
}
#[cfg(test)]
mod tests;