mod batch;
pub mod hooks;
mod savepoint;
mod with_time;
use sqlx::{Acquire, Transaction};
use crate::{clock::ClockHandle, db, one_time_executor::OneTimeExecutor};
pub use batch::*;
pub use savepoint::*;
pub use with_time::*;
pub struct DbOp<'c> {
tx: Transaction<'c, db::Db>,
clock: ClockHandle,
now: Option<chrono::DateTime<chrono::Utc>>,
commit_hooks: Option<hooks::CommitHooks>,
}
impl<'c> DbOp<'c> {
fn new(
tx: Transaction<'c, db::Db>,
clock: ClockHandle,
time: Option<chrono::DateTime<chrono::Utc>>,
) -> Self {
Self {
tx,
clock,
now: time,
commit_hooks: Some(hooks::CommitHooks::new()),
}
}
pub async fn init(pool: &db::Pool) -> Result<DbOp<'static>, sqlx::Error> {
Self::init_with_clock(pool, crate::clock::Clock::handle()).await
}
pub async fn init_with_clock(
pool: &db::Pool,
clock: &ClockHandle,
) -> Result<DbOp<'static>, sqlx::Error> {
let tx = pool.begin().await?;
let time = clock.manual_now();
Ok(DbOp::new(tx, clock.clone(), time))
}
pub fn with_time(self, time: chrono::DateTime<chrono::Utc>) -> DbOpWithTime<'c> {
DbOpWithTime::new(self, time)
}
pub fn with_clock_time(self) -> DbOpWithTime<'c> {
let time = self.now.unwrap_or_else(|| self.clock.now());
DbOpWithTime::new(self, time)
}
pub async fn with_db_time(mut self) -> Result<DbOpWithTime<'c>, sqlx::Error> {
let time = if let Some(time) = self.now {
time
} else if let Some(manual_time) = self.clock.manual_now() {
manual_time
} else {
db::database_now(&mut *self.tx).await?
};
Ok(DbOpWithTime::new(self, time))
}
pub fn maybe_now(&self) -> Option<chrono::DateTime<chrono::Utc>> {
self.now
}
pub async fn begin(&mut self) -> Result<DbOp<'_>, sqlx::Error> {
Ok(DbOp::new(
self.tx.begin().await?,
self.clock.clone(),
self.now,
))
}
pub async fn with_savepoint<T, E, F>(&mut self, f: F) -> Result<Result<T, E>, sqlx::Error>
where
F: AsyncFnOnce(&mut SavepointOp<'_>) -> Result<T, E>,
{
SavepointOperation::with_savepoint(self, f).await
}
pub async fn begin_savepoint(&mut self) -> Result<SavepointOp<'_>, sqlx::Error> {
SavepointOperation::begin_savepoint(self).await
}
pub async fn commit(mut self) -> Result<(), sqlx::Error> {
let commit_hooks = self.commit_hooks.take().expect("no hooks");
match commit_hooks.execute_pre(&mut self).await {
Ok(post_hooks) => match self.tx.commit().await {
Ok(()) => {
post_hooks.execute();
Ok(())
}
Err(error) => {
post_hooks.execute_rollback();
Err(error)
}
},
Err((error, executed)) => {
let _ = self.tx.rollback().await;
executed.execute_rollback();
Err(error)
}
}
}
pub fn tx_mut(&mut self) -> &mut Transaction<'c, db::Db> {
&mut self.tx
}
}
impl<'o> AtomicOperation for DbOp<'o> {
fn maybe_now(&self) -> Option<chrono::DateTime<chrono::Utc>> {
self.maybe_now()
}
fn clock(&self) -> &ClockHandle {
&self.clock
}
fn connection(&mut self) -> &mut db::Connection {
self.tx.connection()
}
fn add_commit_hook_dyn(
&mut self,
type_id: std::any::TypeId,
hook: Box<dyn hooks::DynHook>,
) -> Result<(), Box<dyn hooks::DynHook>> {
self.commit_hooks
.as_mut()
.expect("no hooks")
.push_or_merge(type_id, hook);
Ok(())
}
fn commit_hook_dyn(&self, type_id: std::any::TypeId) -> Option<&dyn hooks::DynHook> {
self.commit_hooks.as_ref()?.get_last_dyn(type_id)
}
fn supports_hooks(&self) -> bool {
true
}
fn savepoint_parts(&mut self) -> (&mut db::Connection, savepoint::HookSlot<'_>) {
(
self.tx.connection(),
savepoint::HookSlot(self.commit_hooks.as_mut()),
)
}
}
pub struct DbOpWithTime<'c> {
inner: DbOp<'c>,
now: chrono::DateTime<chrono::Utc>,
}
impl<'c> DbOpWithTime<'c> {
fn new(mut inner: DbOp<'c>, time: chrono::DateTime<chrono::Utc>) -> Self {
inner.now = Some(time);
Self { inner, now: time }
}
pub fn now(&self) -> chrono::DateTime<chrono::Utc> {
self.now
}
pub async fn begin(&mut self) -> Result<DbOpWithTime<'_>, sqlx::Error> {
Ok(DbOpWithTime::new(self.inner.begin().await?, self.now))
}
pub async fn with_savepoint<T, E, F>(&mut self, f: F) -> Result<Result<T, E>, sqlx::Error>
where
F: AsyncFnOnce(&mut SavepointOp<'_>) -> Result<T, E>,
{
SavepointOperation::with_savepoint(self, f).await
}
pub async fn begin_savepoint(&mut self) -> Result<SavepointOp<'_>, sqlx::Error> {
SavepointOperation::begin_savepoint(self).await
}
pub async fn commit(self) -> Result<(), sqlx::Error> {
self.inner.commit().await
}
pub fn tx_mut(&mut self) -> &mut Transaction<'c, db::Db> {
self.inner.tx_mut()
}
}
impl<'o> AtomicOperation for DbOpWithTime<'o> {
fn maybe_now(&self) -> Option<chrono::DateTime<chrono::Utc>> {
Some(self.now())
}
fn clock(&self) -> &ClockHandle {
self.inner.clock()
}
fn connection(&mut self) -> &mut db::Connection {
self.inner.connection()
}
fn add_commit_hook_dyn(
&mut self,
type_id: std::any::TypeId,
hook: Box<dyn hooks::DynHook>,
) -> Result<(), Box<dyn hooks::DynHook>> {
self.inner.add_commit_hook_dyn(type_id, hook)
}
fn commit_hook_dyn(&self, type_id: std::any::TypeId) -> Option<&dyn hooks::DynHook> {
self.inner.commit_hook_dyn(type_id)
}
fn supports_hooks(&self) -> bool {
self.inner.supports_hooks()
}
fn savepoint_parts(&mut self) -> (&mut db::Connection, savepoint::HookSlot<'_>) {
self.inner.savepoint_parts()
}
}
impl<'o> AtomicOperationWithTime for DbOpWithTime<'o> {
fn now(&self) -> chrono::DateTime<chrono::Utc> {
self.now
}
}
pub trait AtomicOperation: Send {
fn maybe_now(&self) -> Option<chrono::DateTime<chrono::Utc>> {
None
}
fn clock(&self) -> &ClockHandle {
crate::clock::Clock::handle()
}
fn connection(&mut self) -> &mut db::Connection;
fn as_executor(&mut self) -> OneTimeExecutor<'_, &mut db::Connection> {
let now = self.maybe_now();
OneTimeExecutor::new(self.connection(), now)
}
fn add_commit_hook_dyn(
&mut self,
_type_id: std::any::TypeId,
hook: Box<dyn hooks::DynHook>,
) -> Result<(), Box<dyn hooks::DynHook>> {
Err(hook)
}
fn commit_hook_dyn(&self, _type_id: std::any::TypeId) -> Option<&dyn hooks::DynHook> {
None
}
fn add_commit_hook<H: hooks::CommitHook>(&mut self, hook: H) -> Result<(), H>
where
Self: Sized,
{
self.add_commit_hook_dyn(std::any::TypeId::of::<H>(), Box::new(hook))
.map_err(|hook| {
*hook
.into_any()
.downcast::<H>()
.unwrap_or_else(|_| panic!("hook type mismatch"))
})
}
fn commit_hook<H: hooks::CommitHook>(&self) -> Option<&H>
where
Self: Sized,
{
self.commit_hook_dyn(std::any::TypeId::of::<H>())?
.as_any()
.downcast_ref::<H>()
}
fn supports_hooks(&self) -> bool {
false
}
fn savepoint_parts(&mut self) -> (&mut db::Connection, savepoint::HookSlot<'_>) {
(self.connection(), savepoint::HookSlot::unsupported())
}
}
impl<'c> AtomicOperation for sqlx::Transaction<'c, db::Db> {
fn connection(&mut self) -> &mut db::Connection {
&mut *self
}
}
impl<O: AtomicOperation + ?Sized> AtomicOperation for &mut O {
fn maybe_now(&self) -> Option<chrono::DateTime<chrono::Utc>> {
O::maybe_now(&**self)
}
fn clock(&self) -> &ClockHandle {
O::clock(&**self)
}
fn connection(&mut self) -> &mut db::Connection {
O::connection(&mut **self)
}
fn as_executor(&mut self) -> OneTimeExecutor<'_, &mut db::Connection> {
O::as_executor(&mut **self)
}
fn add_commit_hook_dyn(
&mut self,
type_id: std::any::TypeId,
hook: Box<dyn hooks::DynHook>,
) -> Result<(), Box<dyn hooks::DynHook>> {
O::add_commit_hook_dyn(&mut **self, type_id, hook)
}
fn commit_hook_dyn(&self, type_id: std::any::TypeId) -> Option<&dyn hooks::DynHook> {
O::commit_hook_dyn(&**self, type_id)
}
fn supports_hooks(&self) -> bool {
O::supports_hooks(&**self)
}
fn savepoint_parts(&mut self) -> (&mut db::Connection, savepoint::HookSlot<'_>) {
O::savepoint_parts(&mut **self)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn atomic_operation_is_object_safe() {
fn assert_object_safe(_: &mut dyn AtomicOperation) {}
let _ = assert_object_safe as fn(&mut dyn AtomicOperation);
}
}