use std::any::{Any, TypeId};
use std::ops::{Deref, DerefMut};
use std::sync::{Arc, MutexGuard, TryLockError};
use std::time::{Duration, Instant};
use anyhow::{anyhow, Result};
use lora_executor::ExecutorError;
use lora_store::{GraphStorage, GraphStorageMut, InMemoryGraph, MutationRecorder};
use lora_wal::WalRecorder;
use crate::database::Database;
use crate::wal::write_scope::{ensure_wal_query_can_start, WalAbortPolicy, WalWriteScope};
use super::replay::install_recorder_if_inmemory;
pub(crate) struct WriteGuard<'db, S> {
db: &'db Database<S>,
_writer_lock: MutexGuard<'db, ()>,
staged: Option<S>,
}
impl<S> Deref for WriteGuard<'_, S> {
type Target = S;
fn deref(&self) -> &S {
self.staged
.as_ref()
.expect("staged graph already published or taken")
}
}
impl<S> DerefMut for WriteGuard<'_, S> {
fn deref_mut(&mut self) -> &mut S {
self.staged
.as_mut()
.expect("staged graph already published or taken")
}
}
impl<S> WriteGuard<'_, S>
where
S: Send + Sync + 'static,
{
pub(crate) fn publish(mut self) {
if let Some(staged) = self.staged.take() {
self.db.store.store(Arc::new(staged));
}
}
}
impl<S> Database<S>
where
S: GraphStorage + GraphStorageMut + Any + Clone + Send + Sync + 'static,
{
pub(crate) fn write_store(&self) -> WriteGuard<'_, S> {
let lock = self
.writer
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let snapshot = self.store.load_full();
let staged: S = (*snapshot).clone();
WriteGuard {
db: self,
_writer_lock: lock,
staged: Some(staged),
}
}
pub(crate) fn read_store_with_epoch_deadline(
&self,
_deadline: Option<Instant>,
) -> Result<(Arc<S>, u64)> {
Ok(self.store.load_full_with_epoch())
}
pub(crate) fn write_store_deadline(
&self,
deadline: Option<Instant>,
) -> Result<WriteGuard<'_, S>> {
let Some(deadline) = deadline else {
return Ok(self.write_store());
};
loop {
match self.writer.try_lock() {
Ok(lock) => {
let snapshot = self.store.load_full();
let staged: S = (*snapshot).clone();
return Ok(WriteGuard {
db: self,
_writer_lock: lock,
staged: Some(staged),
});
}
Err(TryLockError::Poisoned(poisoned)) => {
let lock = poisoned.into_inner();
let snapshot = self.store.load_full();
let staged: S = (*snapshot).clone();
return Ok(WriteGuard {
db: self,
_writer_lock: lock,
staged: Some(staged),
});
}
Err(TryLockError::WouldBlock) if Instant::now() >= deadline => {
return Err(ExecutorError::QueryTimeout.into());
}
Err(TryLockError::WouldBlock) => {
std::thread::sleep(Duration::from_millis(1));
}
}
}
}
pub(crate) fn observe_snapshot_commit_if_needed(
&self,
store: &S,
recorder: &WalRecorder,
) -> Result<()> {
let Some(snapshots) = &self.snapshots else {
return Ok(());
};
let graph = (store as &dyn Any)
.downcast_ref::<InMemoryGraph>()
.ok_or_else(|| anyhow!("managed snapshots require InMemoryGraph storage"))?;
snapshots.observe_commit(graph, recorder)?;
Ok(())
}
pub(crate) fn with_logged_store_mut<R>(
&self,
f: impl FnOnce(&mut S) -> Result<R>,
) -> Result<R> {
if TypeId::of::<S>() == TypeId::of::<InMemoryGraph>() {
return self.run_with_durable_recorder(f);
}
let guard = self.write_store();
self.with_logged_write_guard(guard, WalAbortPolicy::AbortOnly, f)
}
pub(crate) fn run_with_durable_recorder<R>(
&self,
f: impl FnOnce(&mut S) -> Result<R>,
) -> Result<R> {
let _commit_lock = self
.writer
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(rec) = self.wal.as_ref() {
ensure_wal_query_can_start(rec)?;
rec.arm()?;
}
let mut handle = self.store.write();
let result = {
let staged = handle.as_mut();
if let Some(rec) = self.wal.as_ref() {
install_recorder_if_inmemory(
staged,
Some(rec.clone() as Arc<dyn MutationRecorder>),
);
}
f(staged)
};
if let Some(rec) = self.wal.as_ref() {
match &result {
Ok(_) => {
if rec.commit()?.wrote() {
let live = handle.snapshot();
self.observe_snapshot_commit_if_needed(&*live, rec)?;
}
}
Err(_) => {
let _ = rec.abort();
}
}
}
result
}
pub(crate) fn with_logged_write_guard<R>(
&self,
mut guard: WriteGuard<'_, S>,
abort_policy: WalAbortPolicy,
f: impl FnOnce(&mut S) -> Result<R>,
) -> Result<R> {
let Some(rec) = self.wal.clone() else {
let result = f(&mut *guard);
if result.is_ok() {
guard.publish();
}
return result;
};
install_recorder_if_inmemory(&mut *guard, Some(rec.clone() as Arc<dyn MutationRecorder>));
let scope = WalWriteScope::arm(&rec, abort_policy)?;
let result = f(&mut *guard);
let wrote_commit = scope.finish(&result)?;
if wrote_commit {
self.observe_snapshot_commit_if_needed(&*guard, &rec)?;
}
install_recorder_if_inmemory(&mut *guard, None);
if result.is_ok() {
install_recorder_if_inmemory(
&mut *guard,
Some(rec.clone() as Arc<dyn MutationRecorder>),
);
guard.publish();
}
result
}
}