use crate::db::{
BatchOp, FsyncPolicy, Precondition, WriteAuthz, LOCK_POLL_INTERVAL, WRITE_LOCK_WAIT,
};
use crate::reader::ReaderSnapshot;
use crate::GraphDb;
use core_storage::sync_wal_at;
use core_storage::truncate_wal_at;
use core_storage::GraphError;
use core_storage::RealFs;
use core_storage::Result;
use std::ops::{Deref, DerefMut};
use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Condvar, Mutex, RwLock};
use std::thread;
use std::time::{Duration, Instant};
const MAX_GROUP_SIZE: usize = 256;
type SyncWalFn = Arc<dyn Fn(&Path) -> std::io::Result<()> + Send + Sync>;
struct Submission {
ops: Vec<BatchOp>,
preconds: Vec<Precondition>,
authz_role: Option<String>,
done: std::sync::mpsc::SyncSender<Result<(usize, usize)>>,
}
struct WriteQueue {
pending: Mutex<Vec<Submission>>,
notify: Condvar,
shutdown: AtomicBool,
degraded_msg: Mutex<Option<String>>,
}
impl WriteQueue {
fn new() -> Arc<Self> {
Arc::new(Self {
pending: Mutex::new(Vec::new()),
notify: Condvar::new(),
shutdown: AtomicBool::new(false),
degraded_msg: Mutex::new(None),
})
}
fn enqueue(&self, sub: Submission) {
self.pending
.lock()
.unwrap_or_else(|e| e.into_inner())
.push(sub);
self.notify.notify_one();
}
fn signal_shutdown(&self) {
let _lock = self.pending.lock().unwrap_or_else(|e| e.into_inner());
self.shutdown.store(true, Ordering::Release);
self.notify.notify_all();
}
fn set_degraded(&self, msg: String) {
*self.degraded_msg.lock().unwrap_or_else(|e| e.into_inner()) = Some(msg);
}
fn degraded_message(&self) -> Option<String> {
self.degraded_msg
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone()
}
fn wait_and_drain(&self) -> Vec<Submission> {
let mut lock = self.pending.lock().unwrap_or_else(|e| e.into_inner());
loop {
if !lock.is_empty() {
let n = lock.len().min(MAX_GROUP_SIZE);
return lock.drain(..n).collect();
}
if self.shutdown.load(Ordering::Acquire) {
return vec![];
}
lock = self.notify.wait(lock).unwrap_or_else(|e| e.into_inner());
}
}
}
struct DrainHandle {
queue: Arc<WriteQueue>,
handle: Option<thread::JoinHandle<()>>,
}
impl Drop for DrainHandle {
fn drop(&mut self) {
self.queue.signal_shutdown();
if let Some(h) = self.handle.take() {
let _ = h.join();
}
}
}
pub struct WriteGuard<'a> {
inner: std::sync::RwLockWriteGuard<'a, GraphDb<RealFs>>,
_wal: std::sync::MutexGuard<'a, ()>,
}
impl<'a> Drop for WriteGuard<'a> {
fn drop(&mut self) {
self.inner.end_write_lock();
}
}
impl<'a> Deref for WriteGuard<'a> {
type Target = GraphDb<RealFs>;
fn deref(&self) -> &GraphDb<RealFs> {
&self.inner
}
}
impl<'a> DerefMut for WriteGuard<'a> {
fn deref_mut(&mut self) -> &mut GraphDb<RealFs> {
&mut self.inner
}
}
#[derive(Clone)]
pub struct SharedDb {
inner: Arc<RwLock<GraphDb<RealFs>>>,
queue: Arc<WriteQueue>,
_drain: Arc<DrainHandle>,
wal_mu: Arc<Mutex<()>>,
}
const _: () = {
fn assert_send_sync<T: Send + Sync>() {}
let _ = assert_send_sync::<SharedDb>;
};
impl SharedDb {
pub fn open(dir: &Path) -> Result<Self> {
let db = GraphDb::open_unlocked(dir)?;
Ok(Self::from_db_and_dir_with_sync(
db,
dir.to_path_buf(),
Arc::new(sync_wal_at),
))
}
pub fn open_with_test_sync(
dir: &Path,
sync: impl Fn(&Path) -> std::io::Result<()> + Send + Sync + 'static,
) -> Result<Self> {
let db = GraphDb::open_unlocked(dir)?;
Ok(Self::from_db_and_dir_with_sync(
db,
dir.to_path_buf(),
Arc::new(sync),
))
}
fn from_db_and_dir_with_sync(
db: GraphDb<RealFs>,
dir: std::path::PathBuf,
sync_fn: SyncWalFn,
) -> Self {
let inner = Arc::new(RwLock::new(db));
let queue = WriteQueue::new();
let wal_mu = Arc::new(Mutex::new(()));
let dir_arc = Arc::new(dir);
let drain_inner = Arc::clone(&inner);
let drain_queue = Arc::clone(&queue);
let drain_dir = Arc::clone(&dir_arc);
let drain_wal_mu = Arc::clone(&wal_mu);
let drain_sync_fn = Arc::clone(&sync_fn);
let handle = thread::Builder::new()
.name("groupcommit-drain".into())
.spawn(move || {
drain_loop(
drain_inner,
drain_queue,
drain_dir,
drain_wal_mu,
drain_sync_fn,
)
})
.expect("failed to spawn group-commit drain thread");
SharedDb {
inner,
queue: Arc::clone(&queue),
_drain: Arc::new(DrainHandle {
queue,
handle: Some(handle),
}),
wal_mu,
}
}
fn refresh_if_stale(&self) {
let stale = self
.inner
.read()
.unwrap_or_else(|e| e.into_inner())
.is_stale()
.unwrap_or(false);
if stale {
let mut db = self.inner.write().unwrap_or_else(|e| e.into_inner());
if let Err(e) = db.refresh() {
self.queue.set_degraded(format!(
"refresh failed while following another process's commits: {e}; \
reopen required"
));
}
}
}
pub fn read(&self) -> impl Deref<Target = GraphDb<RealFs>> + '_ {
self.refresh_if_stale();
self.inner.read().unwrap_or_else(|e| e.into_inner())
}
pub fn write(&self) -> WriteGuard<'_> {
let _wal = self.wal_mu.lock().unwrap_or_else(|e| e.into_inner());
let acquired = self
.poll_cross_process_lock(WRITE_LOCK_WAIT)
.unwrap_or(false);
self.enter_write_scope(_wal, acquired).0
}
pub fn write_with_wait(&self, wait: Duration) -> Result<WriteGuard<'_>> {
let _wal = self.wal_mu.lock().unwrap_or_else(|e| e.into_inner());
if !self.poll_cross_process_lock(wait)? {
return Err(GraphError::Busy { holder: None });
}
let (guard, entered) = self.enter_write_scope(_wal, true);
entered?;
Ok(guard)
}
fn poll_cross_process_lock(&self, wait: Duration) -> Result<bool> {
let deadline = Instant::now() + wait;
loop {
{
let db = self.inner.read().unwrap_or_else(|e| e.into_inner());
if db.try_cross_process_lock()? {
return Ok(true);
}
}
let now = Instant::now();
if now >= deadline {
return Ok(false);
}
std::thread::sleep(LOCK_POLL_INTERVAL.min(deadline.saturating_duration_since(now)));
}
}
fn enter_write_scope<'a>(
&'a self,
wal: std::sync::MutexGuard<'a, ()>,
acquired: bool,
) -> (WriteGuard<'a>, Result<()>) {
let mut inner = self.inner.write().unwrap_or_else(|e| e.into_inner());
let entered = inner.enter_write_scope(acquired);
(WriteGuard { inner, _wal: wal }, entered)
}
pub fn reader(&self) -> ReaderSnapshot {
self.read().reader()
}
pub fn submit_batch(&self, ops: Vec<BatchOp>) -> Result<(usize, usize)> {
if let Some(msg) = self.queue.degraded_message() {
return Err(GraphError::Io(std::io::Error::other(msg)));
}
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.queue.enqueue(Submission {
ops,
preconds: Vec::new(),
authz_role: None,
done: tx,
});
rx.recv().unwrap_or_else(|_| {
Err(GraphError::Io(std::io::Error::other(
"group-commit drain thread terminated unexpectedly",
)))
})
}
pub fn submit_batch_cas(
&self,
preconds: Vec<Precondition>,
ops: Vec<BatchOp>,
) -> Result<(usize, usize)> {
if let Some(msg) = self.queue.degraded_message() {
return Err(GraphError::Io(std::io::Error::other(msg)));
}
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.queue.enqueue(Submission {
ops,
preconds,
authz_role: None,
done: tx,
});
rx.recv().unwrap_or_else(|_| {
Err(GraphError::Io(std::io::Error::other(
"group-commit drain thread terminated unexpectedly",
)))
})
}
pub fn submit_batch_authz(&self, role: String, ops: Vec<BatchOp>) -> Result<(usize, usize)> {
if let Some(msg) = self.queue.degraded_message() {
return Err(GraphError::Io(std::io::Error::other(msg)));
}
let (tx, rx) = std::sync::mpsc::sync_channel(1);
self.queue.enqueue(Submission {
ops,
preconds: Vec::new(),
authz_role: Some(role),
done: tx,
});
rx.recv().unwrap_or_else(|_| {
Err(GraphError::Io(std::io::Error::other(
"group-commit drain thread terminated unexpectedly",
)))
})
}
}
fn poll_cross_process_lock(inner: &Arc<RwLock<GraphDb<RealFs>>>, wait: Duration) -> Result<bool> {
let deadline = Instant::now() + wait;
loop {
{
let db = inner.read().unwrap_or_else(|e| e.into_inner());
if db.try_cross_process_lock()? {
return Ok(true);
}
}
let now = Instant::now();
if now >= deadline {
return Ok(false);
}
std::thread::sleep(LOCK_POLL_INTERVAL.min(deadline.saturating_duration_since(now)));
}
}
fn drain_loop(
inner: Arc<RwLock<GraphDb<RealFs>>>,
queue: Arc<WriteQueue>,
dir: Arc<std::path::PathBuf>,
wal_mu: Arc<Mutex<()>>,
sync_fn: SyncWalFn,
) {
loop {
let mut group = queue.wait_and_drain();
if group.is_empty() {
return; }
let submissions: Vec<(Vec<Precondition>, Vec<BatchOp>, Option<String>)> = group
.iter_mut()
.map(|s| {
(
std::mem::take(&mut s.preconds),
std::mem::take(&mut s.ops),
s.authz_role.take(),
)
})
.collect();
let wal_guard = wal_mu.lock().unwrap_or_else(|e| e.into_inner());
let lock_outcome = poll_cross_process_lock(&inner, WRITE_LOCK_WAIT).and_then(|acquired| {
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
db.enter_write_scope(acquired).map(|()| acquired)
});
match lock_outcome {
Ok(true) => {}
Ok(false) => {
{
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
db.end_write_lock();
}
drop(wal_guard);
for sub in group {
let _ = sub.done.send(Err(GraphError::Busy { holder: None }));
}
continue;
}
Err(e) => {
let reason = format!("refresh under the write lock failed: {e}; reopen required");
{
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
db.discard_deferred_events();
db.set_deferred_events_mode(false);
db.set_degraded();
db.end_write_lock();
}
queue.set_degraded(reason.clone());
drop(wal_guard);
for sub in group {
let _ = sub
.done
.send(Err(GraphError::Io(std::io::Error::other(reason.clone()))));
}
return;
}
}
let pre_group_wal_len = std::fs::metadata(dir.join("wal.bin"))
.map(|m| m.len())
.unwrap_or(0);
let (results, should_sync): (Vec<Result<(usize, usize)>>, bool) = {
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
let sync_needed = db.fsync_policy() != FsyncPolicy::Relaxed;
if sync_needed {
db.set_deferred_events_mode(true);
}
let mut r: Vec<Result<(usize, usize)>> = Vec::with_capacity(submissions.len());
let mut pending_non_cas: Vec<Vec<BatchOp>> = Vec::new();
for (preconds, ops, authz_role) in submissions {
if let Some(role) = authz_role {
if !pending_non_cas.is_empty() {
let batch = std::mem::take(&mut pending_non_cas);
r.extend(db.commit_group_nosync(batch));
}
let result = (|| -> Result<(usize, usize)> {
let scope = {
let roles_vec = db.roles();
let def = roles_vec
.iter()
.find(|r| r.name == role)
.ok_or_else(|| GraphError::KeyNotFound {
key: format!("role:{role}"),
})?
.clone();
def.write.ok_or_else(|| GraphError::RoleWriteDenied {
reason: "role-bound token: writes are not permitted".into(),
})?
};
let mask = db.mask_for_role(&role)?;
let authz = WriteAuthz { role, scope, mask };
db.write_batch_authz_nosync(Some(&authz), ops)
})();
r.push(result);
} else if preconds.is_empty() {
pending_non_cas.push(ops);
} else {
if !pending_non_cas.is_empty() {
let batch = std::mem::take(&mut pending_non_cas);
r.extend(db.commit_group_nosync(batch));
}
let result = match db.check_preconditions(&preconds) {
Ok(()) => db
.commit_group_nosync(vec![ops])
.into_iter()
.next()
.unwrap_or(Ok((0, 0))),
Err(e) => Err(e),
};
r.push(result);
}
}
if !pending_non_cas.is_empty() {
r.extend(db.commit_group_nosync(pending_non_cas));
}
(r, sync_needed)
};
let sync_result: Result<()> = if should_sync && results.iter().any(|r| r.is_ok()) {
sync_fn(&dir).map_err(GraphError::Io)
} else {
Ok(()) };
if let Err(ref io_err) = sync_result {
if results.iter().any(|r| r.is_ok()) {
let _ = truncate_wal_at(&dir, pre_group_wal_len);
}
{
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
db.discard_deferred_events();
db.set_deferred_events_mode(false);
db.set_degraded();
db.set_wal_consumed(pre_group_wal_len);
db.end_write_lock();
}
queue.set_degraded(io_err.to_string());
drop(wal_guard);
for sub in group {
let _ = sub.done.send(Err(GraphError::Io(std::io::Error::other(
"group-commit fsync failed; database is degraded, reopen required",
))));
}
return; }
if should_sync {
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
db.flush_deferred_events();
db.set_deferred_events_mode(false);
}
{
let mut db = inner.write().unwrap_or_else(|e| e.into_inner());
db.end_write_lock();
}
drop(wal_guard);
for (sub, result) in group.into_iter().zip(results) {
let _ = sub.done.send(result);
}
}
}