use async_trait::async_trait;
use crate::dialect::Dialect;
#[derive(Debug, Clone)]
pub struct Row {
columns: Vec<String>,
values: Vec<serde_json::Value>,
}
impl Row {
pub fn new(columns: Vec<String>, values: Vec<serde_json::Value>) -> Self {
debug_assert_eq!(
columns.len(),
values.len(),
"columns and values must have the same length"
);
Self { columns, values }
}
pub fn get(&self, col: &str) -> Option<&serde_json::Value> {
self.columns
.iter()
.position(|c| c == col)
.and_then(|i| self.values.get(i))
}
}
#[async_trait]
pub trait TransactionInner: Send {
async fn execute(&mut self, sql: &str) -> Result<u64, RdbmsError>;
async fn query(&mut self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
async fn execute_with(
&mut self,
sql: &str,
params: &[serde_json::Value],
) -> Result<u64, RdbmsError>;
async fn query_with(
&mut self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError>;
fn dialect(&self) -> Dialect;
async fn commit(&mut self) -> Result<(), RdbmsError>;
async fn rollback(&mut self) -> Result<(), RdbmsError>;
}
const NO_BACKING: &str = "transaction has no backing connection (created via Transaction::new)";
pub struct Transaction {
committed: bool,
rolled_back: bool,
dialect: Dialect,
inner: tokio::sync::Mutex<Option<Box<dyn TransactionInner>>>,
}
impl Transaction {
pub fn new() -> Self {
Self {
committed: false,
rolled_back: false,
dialect: Dialect::Standard,
inner: tokio::sync::Mutex::new(None),
}
}
pub fn with_inner(inner: Box<dyn TransactionInner>) -> Self {
let dialect = inner.dialect();
Self {
committed: false,
rolled_back: false,
dialect,
inner: tokio::sync::Mutex::new(Some(inner)),
}
}
pub async fn commit(mut self) -> Result<(), RdbmsError> {
if let Some(inner) = self.inner.get_mut().as_mut() {
inner.commit().await?;
}
self.committed = true;
Ok(())
}
pub async fn rollback(mut self) -> Result<(), RdbmsError> {
if let Some(inner) = self.inner.get_mut().as_mut() {
inner.rollback().await?;
}
self.committed = false;
self.rolled_back = true;
Ok(())
}
}
impl Default for Transaction {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl SqlExecutor for Transaction {
async fn execute(&self, sql: &str) -> Result<u64, RdbmsError> {
let mut guard = self.inner.lock().await;
match guard.as_mut() {
Some(inner) => inner.execute(sql).await,
None => Err(RdbmsError::Database(NO_BACKING.into())),
}
}
async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError> {
let mut guard = self.inner.lock().await;
match guard.as_mut() {
Some(inner) => inner.query(sql).await,
None => Err(RdbmsError::Database(NO_BACKING.into())),
}
}
async fn execute_with(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<u64, RdbmsError> {
let mut guard = self.inner.lock().await;
match guard.as_mut() {
Some(inner) => inner.execute_with(sql, params).await,
None => Err(RdbmsError::Database(NO_BACKING.into())),
}
}
async fn query_with(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
let mut guard = self.inner.lock().await;
match guard.as_mut() {
Some(inner) => inner.query_with(sql, params).await,
None => Err(RdbmsError::Database(NO_BACKING.into())),
}
}
async fn execute_then_query(
&self,
first: &str,
first_params: &[serde_json::Value],
second: &str,
) -> Result<Vec<Row>, RdbmsError> {
let mut guard = self.inner.lock().await;
match guard.as_mut() {
Some(inner) => {
inner.execute_with(first, first_params).await?;
inner.query(second).await
}
None => Err(RdbmsError::Database(NO_BACKING.into())),
}
}
fn dialect(&self) -> Dialect {
self.dialect
}
}
impl Drop for Transaction {
fn drop(&mut self) {
if !self.committed && !self.rolled_back {
crate::timeout::TRANSACTIONS_LEAKED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
tracing::warn!("transaction dropped without commit — rolling back");
}
}
}
#[async_trait]
pub trait SqlExecutor: Send + Sync {
async fn execute(&self, sql: &str) -> Result<u64, RdbmsError>;
async fn query(&self, sql: &str) -> Result<Vec<Row>, RdbmsError>;
async fn execute_with(
&self,
_sql: &str,
_params: &[serde_json::Value],
) -> Result<u64, RdbmsError> {
Err(RdbmsError::Database(
"parameterized execute not supported by this backend".into(),
))
}
async fn query_with(
&self,
_sql: &str,
_params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
Err(RdbmsError::Database(
"parameterized query not supported by this backend".into(),
))
}
async fn query_write(
&self,
sql: &str,
params: &[serde_json::Value],
) -> Result<Vec<Row>, RdbmsError> {
self.query_with(sql, params).await
}
async fn execute_then_query(
&self,
_first: &str,
_first_params: &[serde_json::Value],
_second: &str,
) -> Result<Vec<Row>, RdbmsError> {
Err(RdbmsError::Database(
"this backend cannot run two statements atomically on one connection".into(),
))
}
fn dialect(&self) -> Dialect;
}
#[async_trait]
pub trait RdbmsClient: SqlExecutor {
async fn transaction(&self) -> Result<Transaction, RdbmsError>;
}
#[derive(Debug, thiserror::Error)]
pub enum RdbmsError {
#[error("database error: {0}")]
Database(String),
#[error("connection error: {0}")]
Connection(String),
#[error("configuration error: {0}")]
Config(String),
#[error("timeout: {0}")]
Timeout(String),
#[error("no available replica")]
NoAvailableReplica,
}
#[cfg(test)]
mod tests;