use std::any::{Any, TypeId};
use std::ops::{Deref, DerefMut};
use std::sync::{Arc, Mutex, MutexGuard, TryLockError};
use std::time::{Duration, Instant};
use anyhow::{anyhow, Result};
use lora_executor::ExecutorError;
use lora_store::{GraphStorage, GraphStorageMut, InMemoryGraph, MutationEvent, MutationRecorder};
use lora_wal::WalRecorder;
use crate::database::Database;
use crate::transaction::BufferingRecorder;
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_deadline(&self, _deadline: Option<Instant>) -> Result<Arc<S>> {
Ok(self.store.load_full())
}
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.with_live_store_mut(f);
}
let guard = self.write_store();
self.with_logged_write_guard(guard, WalAbortPolicy::AbortOnly, f)
}
fn with_live_store_mut<R>(&self, f: impl FnOnce(&mut S) -> Result<R>) -> Result<R> {
let _commit_lock = self
.writer
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let buffer = Arc::new(Mutex::new(Vec::<MutationEvent>::new()));
let buffering_rec: Arc<dyn MutationRecorder> =
Arc::new(BufferingRecorder::new(buffer.clone()));
let mut handle = self.store.write();
let exec_result = {
let staged = handle.as_mut();
install_recorder_if_inmemory(staged, Some(buffering_rec));
let r = f(staged);
install_recorder_if_inmemory(staged, None);
r
};
let events: Vec<MutationEvent> = std::mem::take(&mut buffer.lock().unwrap());
let mut wrote_commit = false;
if let Some(rec) = self.wal.as_ref() {
if exec_result.is_ok() && !events.is_empty() {
ensure_wal_query_can_start(rec)?;
wrote_commit = rec.commit_events(events)?.wrote();
}
let staged = handle.as_mut();
install_recorder_if_inmemory(staged, Some(rec.clone() as Arc<dyn MutationRecorder>));
}
if wrote_commit {
if let Some(rec) = self.wal.as_ref() {
let live = handle.snapshot();
self.observe_snapshot_commit_if_needed(&*live, rec)?;
}
}
exec_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
}
}