1use std::sync::Arc;
3use std::sync::atomic::{AtomicUsize, Ordering};
4
5use async_trait::async_trait;
6use ecat_circuit_breaker::{Breaker, BreakerConfig, BreakerState};
7
8use crate::breaker::map_breaker_error;
9use crate::dialect::Dialect;
10use crate::rdbms::{RdbmsClient, RdbmsError, Row, SqlExecutor, Transaction};
11
12struct Endpoint {
18 client: Arc<dyn RdbmsClient>,
19 breaker: Breaker,
20}
21
22impl Endpoint {
23 fn new(client: Arc<dyn RdbmsClient>, cfg: &BreakerConfig) -> Self {
24 Self {
25 client,
26 breaker: Breaker::new(cfg.clone()),
27 }
28 }
29
30 fn is_available(&self) -> bool {
32 self.breaker.state() != BreakerState::Open
33 }
34
35 fn dialect(&self) -> Dialect {
36 self.client.dialect()
37 }
38
39 async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
40 self.breaker
41 .call(|| self.client.execute(sql))
42 .await
43 .map_err(map_breaker_error)
44 }
45
46 async fn execute_with(
47 &self,
48 sql: &str,
49 params: &[serde_json::Value],
50 ) -> Result<u64, RdbmsError> {
51 self.breaker
52 .call(|| self.client.execute_with(sql, params))
53 .await
54 .map_err(map_breaker_error)
55 }
56
57 async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
58 self.breaker
59 .call(|| self.client.query(sql))
60 .await
61 .map_err(map_breaker_error)
62 }
63
64 async fn query_with(
65 &self,
66 sql: &str,
67 params: &[serde_json::Value],
68 ) -> Result<Vec<Row>, RdbmsError> {
69 self.breaker
70 .call(|| self.client.query_with(sql, params))
71 .await
72 .map_err(map_breaker_error)
73 }
74
75 async fn query_write(
76 &self,
77 sql: &str,
78 params: &[serde_json::Value],
79 ) -> Result<Vec<Row>, RdbmsError> {
80 self.breaker
81 .call(|| self.client.query_write(sql, params))
82 .await
83 .map_err(map_breaker_error)
84 }
85
86 async fn execute_then_query(
87 &self,
88 first: &str,
89 first_params: &[serde_json::Value],
90 second: &str,
91 ) -> Result<Vec<Row>, RdbmsError> {
92 self.breaker
93 .call(|| self.client.execute_then_query(first, first_params, second))
94 .await
95 .map_err(map_breaker_error)
96 }
97}
98
99pub struct RdbmsRouting {
104 primary: Endpoint,
105 replicas: Vec<Endpoint>,
106 next: AtomicUsize,
108 fallback_to_primary: bool,
110}
111
112impl RdbmsRouting {
113 pub fn new(primary: Arc<dyn RdbmsClient>, replicas: Vec<Arc<dyn RdbmsClient>>) -> Self {
115 Self::with_breaker_config(primary, replicas, BreakerConfig::default())
116 }
117
118 pub fn with_breaker_config(
120 primary: Arc<dyn RdbmsClient>,
121 replicas: Vec<Arc<dyn RdbmsClient>>,
122 cfg: BreakerConfig,
123 ) -> Self {
124 Self {
125 primary: Endpoint::new(primary, &cfg),
126 replicas: replicas
127 .into_iter()
128 .map(|client| Endpoint::new(client, &cfg))
129 .collect(),
130 next: AtomicUsize::new(0),
131 fallback_to_primary: true,
132 }
133 }
134
135 pub fn fallback_to_primary(mut self, yes: bool) -> Self {
138 self.fallback_to_primary = yes;
139 self
140 }
141
142 fn pick_replica(&self) -> Option<&Endpoint> {
150 let n = self.replicas.len();
151 if n == 0 {
152 return None;
153 }
154 let start = self.next.fetch_add(1, Ordering::Relaxed);
155 (0..n)
156 .map(|i| &self.replicas[(start + i) % n])
157 .find(|ep| ep.is_available())
158 }
159}
160
161#[async_trait]
162impl SqlExecutor for RdbmsRouting {
163 async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
165 self.primary.execute(sql).await
166 }
167
168 async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
170 match self.pick_replica() {
171 Some(replica) => replica.query(sql).await,
172 None if self.fallback_to_primary => self.primary.query(sql).await,
173 None => Err(RdbmsError::NoAvailableReplica),
174 }
175 }
176
177 async fn execute_with(
179 &self,
180 sql: &str,
181 params: &[serde_json::Value],
182 ) -> Result<u64, RdbmsError> {
183 self.primary.execute_with(sql, params).await
184 }
185
186 async fn query_with(
188 &self,
189 sql: &str,
190 params: &[serde_json::Value],
191 ) -> Result<Vec<Row>, RdbmsError> {
192 match self.pick_replica() {
193 Some(replica) => replica.query_with(sql, params).await,
194 None if self.fallback_to_primary => self.primary.query_with(sql, params).await,
195 None => Err(RdbmsError::NoAvailableReplica),
196 }
197 }
198
199 async fn query_write(
202 &self,
203 sql: &str,
204 params: &[serde_json::Value],
205 ) -> Result<Vec<Row>, RdbmsError> {
206 self.primary.query_write(sql, params).await
207 }
208
209 async fn execute_then_query(
211 &self,
212 first: &str,
213 first_params: &[serde_json::Value],
214 second: &str,
215 ) -> Result<Vec<Row>, RdbmsError> {
216 self.primary
217 .execute_then_query(first, first_params, second)
218 .await
219 }
220
221 fn dialect(&self) -> Dialect {
223 self.primary.dialect()
224 }
225}
226
227#[async_trait]
228impl RdbmsClient for RdbmsRouting {
229 async fn transaction(&self) -> Result<Transaction, RdbmsError> {
232 self.primary.client.transaction().await
233 }
234}
235
236#[cfg(test)]
237mod tests;