use std::{
fs, io,
path::{Path, PathBuf},
sync::{Arc, atomic::Ordering::Relaxed},
};
use async_lock::RwLock;
use gxhash::HashMap as GxHashMap;
use papaya::HashMap as PapayaMap;
use parking_lot::Mutex;
use waof::WalLog;
use wdev::Device;
use wkv::WedbStore;
use super::{
database_manager_base::DatabaseManagerBase,
garnet_database::GarnetDatabase,
i_database_manager::{HybridLogStats, IDatabaseManager},
};
use crate::storage::functions::functions_state::FunctionsState;
pub struct MultiDatabaseManager<D: Device> {
pub base: DatabaseManagerBase<D>,
pub store: Arc<WedbStore<D>>,
pub databases: PapayaMap<i64, Arc<GarnetDatabase<D>>>,
pub content_lock: RwLock<()>,
pub wal_factory: Mutex<Option<(Arc<D>, waof::WalConfig)>>,
pub checkpoint_root: PathBuf,
}
impl<D: Device> MultiDatabaseManager<D> {
pub fn new(store: Arc<WedbStore<D>>, checkpoint_root: PathBuf) -> Self {
Self {
base: DatabaseManagerBase::new(checkpoint_root.join("0")),
store,
databases: PapayaMap::new(),
content_lock: RwLock::new(()),
wal_factory: Mutex::new(None),
checkpoint_root,
}
}
pub fn enable_aof(&self, device: Arc<D>, config: waof::WalConfig) {
*self.wal_factory.lock() = Some((device, config));
}
pub fn checkpoint_dir_of(&self, db_id: i64) -> PathBuf {
self.checkpoint_root.join(db_id.to_string())
}
pub fn try_add_database(
&self,
db_id: i64,
store: Arc<WedbStore<D>>,
) -> wkv::Result<Arc<GarnetDatabase<D>>> {
let db = Arc::new(GarnetDatabase::new(
db_id,
store,
Arc::clone(&self.store.device),
self.checkpoint_dir_of(db_id),
self.create_aof(db_id),
));
self.databases.pin().insert(db_id, Arc::clone(&db));
Ok(db)
}
fn create_aof(&self, _db_id: i64) -> Option<Arc<WalLog<D>>> {
let (device, config) = self.wal_factory.lock().clone()?;
WalLog::new(device, config).map(Arc::new).ok()
}
pub fn handle_database_added(&self, db: &GarnetDatabase<D>) {
db.last_save_store_tail_address
.store(db.store.tail_address(), Relaxed);
}
pub fn copy_databases(&self) -> GxHashMap<i64, Arc<GarnetDatabase<D>>> {
self
.databases
.pin()
.iter()
.fold(GxHashMap::default(), |mut m, (k, v)| {
m.insert(*k, Arc::clone(v));
m
})
}
pub fn try_get_databases_content_write_lock(
&self,
) -> Option<async_lock::RwLockWriteGuard<'_, ()>> {
self.content_lock.try_write()
}
pub fn try_get_databases_content_read_lock(&self) -> Option<async_lock::RwLockReadGuard<'_, ()>> {
self.content_lock.try_read()
}
pub fn try_get_saved_database_ids(&self) -> wkv::Result<Vec<i64>> {
let mut ids = Vec::new();
let entries = match fs::read_dir(&self.checkpoint_root) {
Ok(entries) => entries,
Err(e) if e.kind() == io::ErrorKind::NotFound => {
log::debug!(
"检查点根目录不存在,按空集处理: checkpoint_root = {}",
self.checkpoint_root.display()
);
return Ok(ids);
}
Err(e) => {
log::error!(
"枚举已持久化库编号失败: checkpoint_root = {}; err = {e}",
self.checkpoint_root.display()
);
return Err(wkv::Error::from(e));
}
};
for entry in entries.flatten() {
if entry.path().is_dir()
&& let Ok(n) = entry.file_name().into_string()
&& let Ok(id) = n.parse::<i64>()
{
ids.push(id);
}
}
ids.sort_unstable();
Ok(ids)
}
pub async fn run_paused_checkpoints_and_release_locks(&self) -> wkv::Result<()> {
let pin = self.databases.pin();
for (_, db) in pin.iter() {
self.base.resume_checkpoints(db);
}
Ok(())
}
pub async fn take_one_checkpoint(&self, db: &GarnetDatabase<D>) -> wkv::Result<bool> {
self.base.take_database_checkpoint_async(db).await
}
pub fn update_last_save_data(&self, db: &GarnetDatabase<D>, now_ms: u64) {
db.update_last_save(now_ms);
}
pub fn get_db_by_id(&self, db_id: i64) -> Option<Arc<GarnetDatabase<D>>> {
self.databases.pin().get(&db_id).cloned()
}
async fn collect_db_kv(&self, db_id: i64) -> Vec<(Vec<u8>, Vec<u8>)> {
let Ok(session) = self.store.new_session() else {
return Vec::new();
};
session.set_active_db(db_id.max(0) as u64);
let batch = session.enter_batch();
let ss = crate::storage::session::storage_session::StorageSession::new(batch);
let Ok((map, keys)) = ss.string_snapshot().await else {
return Vec::new();
};
keys
.into_iter()
.filter_map(|k| {
let v = map.get(&k).cloned().flatten()?;
Some((k, v))
})
.collect()
}
async fn delete_db_keys(&self, db_id: i64) {
let Ok(session) = self.store.new_session() else {
return;
};
session.set_active_db(db_id.max(0) as u64);
let batch = session.enter_batch();
let ss = crate::storage::session::storage_session::StorageSession::new(batch);
if let Ok((_, keys)) = ss.string_snapshot().await {
for key in keys {
let _ = ss.delete_string(&key).await;
}
}
}
pub async fn await_commit(&self, db_id: i64) -> wkv::Result<bool> {
self.wait_for_commit_to_aof_async(db_id)
}
pub fn open_database_store(&self, _db_id: i64, _dir: &Path) -> wkv::Result<Arc<WedbStore<D>>> {
Ok(Arc::clone(&self.store))
}
}
impl<D: Device> IDatabaseManager<D> for MultiDatabaseManager<D> {
async fn try_get_or_add_database(
&self,
db_id: i64,
) -> wkv::Result<(Arc<GarnetDatabase<D>>, bool)> {
{
let pin = self.databases.pin();
if let Some(db) = pin.get(&db_id) {
return Ok((Arc::clone(db), false));
}
}
let _guard = self.content_lock.write().await;
let pin = self.databases.pin();
if let Some(db) = pin.get(&db_id) {
return Ok((Arc::clone(db), false));
}
drop(pin);
let db = self.try_add_database(db_id, Arc::clone(&self.store))?;
self.handle_database_added(&db);
Ok((db, true))
}
fn try_get_database(&self, db_id: i64) -> Option<Arc<GarnetDatabase<D>>> {
self.get_db_by_id(db_id)
}
fn try_pause_checkpoints(&self, db_id: i64) -> bool {
self
.get_db_by_id(db_id)
.is_some_and(|db| self.base.try_pause_checkpoints(&db))
}
fn resume_checkpoints(&self, db_id: i64) {
if let Some(db) = self.get_db_by_id(db_id) {
self.base.resume_checkpoints(&db);
}
}
async fn recover_checkpoint_async(
&self,
replica_recover: bool,
recover_from_token: Option<u128>,
) -> wkv::Result<()> {
let _ = replica_recover;
for db_id in self.try_get_saved_database_ids()? {
let (db, _) = self.try_get_or_add_database(db_id).await?;
if let Some(token) = recover_from_token.or_else(|| {
wkv::CheckpointManager::<D>::find_latest_checkpoint(&db.checkpoint_dir)
.ok()
.flatten()
}) {
wkv::CheckpointManager::recover(&db.checkpoint_dir, token, Arc::clone(&db.device))
.await
.map_err(wkv::Error::from)?;
}
self.base.recover_database_aof_async(&db).await?;
}
Ok(())
}
async fn take_checkpoint_async(&self, _background: bool, db_id: i64) -> wkv::Result<bool> {
let mut taken = false;
let targets: Vec<Arc<GarnetDatabase<D>>> = if db_id < 0 {
self.get_databases_snapshot()
} else {
self.get_db_by_id(db_id).into_iter().collect()
};
for db in targets {
taken |= self.take_one_checkpoint(&db).await?;
self.base.run_post_checkpoint_cleanup(&db)?;
}
Ok(taken)
}
async fn take_on_demand_checkpoint_async(&self, entry_ms: u64, db_id: i64) -> wkv::Result<()> {
if let Some(db) = self.get_db_by_id(db_id) {
self
.base
.take_on_demand_checkpoint_async(&db, entry_ms)
.await?;
}
Ok(())
}
async fn task_checkpoint_based_on_aof_size_limit_async(
&self,
aof_size_limit: u64,
) -> wkv::Result<()> {
for (_, db) in self.databases.pin().iter() {
self
.base
.checkpoint_if_aof_exceeds(db, aof_size_limit)
.await?;
}
Ok(())
}
fn commit_to_aof_async(&self, db_id: i64) -> wkv::Result<()> {
if db_id < 0 {
for (_, db) in self.databases.pin().iter() {
self.base.commit_aof(db)?;
}
} else if let Some(db) = self.get_db_by_id(db_id) {
self.base.commit_aof(&db)?;
}
Ok(())
}
fn wait_for_commit_to_aof_async(&self, db_id: i64) -> wkv::Result<bool> {
let ok = self
.get_databases_snapshot()
.into_iter()
.filter(|db| db_id < 0 || db.id == db_id)
.all(|db| {
db.aof
.as_ref()
.is_none_or(|aof| aof.flushed_until_address() >= aof.committed_until_address())
});
Ok(ok)
}
async fn recover_aof_async(&self) -> wkv::Result<u64> {
let mut total = 0u64;
for (_, db) in self.databases.pin().iter() {
total += self.base.recover_database_aof_async(db).await?;
}
Ok(total)
}
async fn replay_aof(&self, until: u64) -> wkv::Result<u64> {
let mut total = 0u64;
for (_, db) in self.databases.pin().iter() {
total += self.base.replay_database_aof(db, 0, until).await?;
}
Ok(total)
}
fn grow_indexes_if_needed_async(&self) -> wkv::Result<bool> {
self.databases.pin().iter().try_fold(false, |acc, (_, db)| {
self.base.grow_index_if_needed_async(db).map(|g| acc | g)
})
}
async fn execute_object_collection(&self, db_id: i64) -> wkv::Result<usize> {
let mut n = 0usize;
for (_, db) in self.databases.pin().iter() {
if db_id < 0 || db.id == db_id {
let session = self.store.new_session()?;
session.set_active_db(db.id.max(0) as u64);
let batch = session.enter_batch();
let storage = crate::storage::session::storage_session::StorageSession::new(batch);
n += storage.object_collect(|_, _| true).await?;
}
}
Ok(n)
}
fn start_size_trackers(&self) {
for (_, db) in self.databases.pin().iter() {
db.size_tracker.restart();
}
}
fn reset_revivification_stats(&self) {
}
fn enqueue_commit(&self, db_id: i64, until: u64) {
if let Some(db) = self.get_db_by_id(db_id) {
db.last_save_store_tail_address.store(until, Relaxed);
}
}
fn get_databases_snapshot(&self) -> Vec<Arc<GarnetDatabase<D>>> {
self
.databases
.pin()
.iter()
.map(|(_, v)| Arc::clone(v))
.collect()
}
async fn flush_database(&self, db_id: i64) -> wkv::Result<()> {
if let Some(db) = self.get_db_by_id(db_id) {
self.base.reset_database(&db).await?;
}
Ok(())
}
async fn flush_all_databases(&self) -> wkv::Result<()> {
let dbs = self.get_databases_snapshot();
for db in dbs {
self.base.reset_database(&db).await?;
}
Ok(())
}
async fn try_swap_databases(&self, db_id1: i64, db_id2: i64) -> bool {
swap_impl::swap(self, db_id1, db_id2).await
}
fn create_functions_state(&self, _db_id: i64) -> FunctionsState {
FunctionsState::new()
}
async fn collect_hybrid_log_stats(&self) -> wkv::Result<Vec<(i64, HybridLogStats)>> {
let mut out = Vec::new();
for (_, db) in self.databases.pin().iter() {
let stats = self.base.collect_hybrid_log_stats_for_db(db).await?;
out.push((db.id, stats));
}
out.sort_by_key(|(id, _)| *id);
Ok(out)
}
fn recover_vector_sets(&self, _db_id: i64) -> wkv::Result<u64> {
Ok(0)
}
}
mod swap_impl {
use std::sync::Arc;
use wdev::Device;
use super::MultiDatabaseManager;
pub(super) async fn swap<D: Device>(
manager: &MultiDatabaseManager<D>,
db_id1: i64,
db_id2: i64,
) -> bool {
if db_id1 == db_id2 {
return true;
}
let db1 = manager.get_db_by_id(db_id1);
let db2 = manager.get_db_by_id(db_id2);
let (Some(db1), Some(db2)) = (db1, db2) else {
return false;
};
let _guard = manager.content_lock.write().await;
let snap1 = manager.collect_db_kv(db_id1).await;
let snap2 = manager.collect_db_kv(db_id2).await;
manager.delete_db_keys(db_id1).await;
manager.delete_db_keys(db_id2).await;
let session = match manager.store.new_session() {
Ok(s) => s,
Err(_) => return false,
};
for (db_id, pairs) in [(db_id2, snap1), (db_id1, snap2)] {
session.set_active_db(db_id.max(0) as u64);
for (key, value) in pairs {
if session.upsert(&key, &value).await.is_err() {
return false;
}
}
}
let pin = manager.databases.pin();
pin.insert(db_id1, Arc::clone(&db2));
pin.insert(db_id2, Arc::clone(&db1));
drop(pin);
let _ = (db1, db2);
true
}
}