use std::any::{Any, TypeId};
use std::sync::{Arc, MutexGuard, TryLockError};
use std::time::Duration;
use web_time::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> 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));
}
}
pub(crate) fn staged_or_error(&self) -> Result<&S> {
self.staged
.as_ref()
.ok_or_else(|| anyhow!("staged graph already published or taken"))
}
pub(crate) fn staged_mut_or_error(&mut self) -> Result<&mut S> {
self.staged
.as_mut()
.ok_or_else(|| anyhow!("staged graph already published or taken"))
}
}
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 guard = self.write_store();
self.with_logged_write_guard(guard, WalAbortPolicy::AbortOnly, f)
}
pub(crate) fn run_live_fast_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 live = handle.as_mut();
if let Some(rec) = self.wal.as_ref() {
install_recorder_if_inmemory(live, Some(rec.clone() as Arc<dyn MutationRecorder>));
}
f(live)
};
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 mutated = rec.abort().unwrap_or(false);
if mutated {
rec.poison(crate::database::QUERY_FAILURE_POISON);
}
}
}
}
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 staged = guard.staged_mut_or_error()?;
let result = f(staged);
if result.is_ok() {
guard.publish();
}
return result;
};
{
let staged = guard.staged_mut_or_error()?;
install_recorder_if_inmemory(staged, Some(rec.clone() as Arc<dyn MutationRecorder>));
}
let scope = WalWriteScope::arm(&rec, abort_policy)?;
let result = {
let staged = guard.staged_mut_or_error()?;
f(staged)
};
let wrote_commit = scope.finish(&result)?;
if wrote_commit {
let staged = guard.staged_or_error()?;
self.observe_snapshot_commit_if_needed(staged, &rec)?;
}
{
let staged = guard.staged_mut_or_error()?;
install_recorder_if_inmemory(staged, None);
}
if result.is_ok() {
{
let staged = guard.staged_mut_or_error()?;
install_recorder_if_inmemory(
staged,
Some(rec.clone() as Arc<dyn MutationRecorder>),
);
}
guard.publish();
}
result
}
}