use std::cell::RefCell;
use std::collections::BTreeMap;
use std::fs::File;
use std::path::{Path, PathBuf};
#[cfg(feature = "counters")]
use std::sync::MutexGuard;
use std::sync::{Arc, Mutex};
#[cfg(not(feature = "counters"))]
use std::sync::{RwLock, RwLockReadGuard, RwLockWriteGuard};
use memmap2::Mmap;
use plugmem_core::{
Config, Error, FactFault, FactRecord, LinkInput, MaintainReport, MemStorage, Memory,
OpenReport, RecallQuery, RecallResult, RecallScratch, RememberInput, RememberOutcome, Stats,
Storage,
};
thread_local! {
static RECALL_SCRATCH: RefCell<RecallScratch> = RefCell::new(RecallScratch::new());
}
use crate::embedder::Embedder;
use crate::error::HostError;
use crate::readonly::ReadOnlyDatabase;
use crate::storage::{FileScratch, FileStorage, FsyncPolicy};
self_cell::self_cell!(
struct OverlayMap {
owner: Mmap,
#[covariant]
dependent: OverlayMemory,
}
);
type OverlayMemory<'a> = Memory<'a>;
enum Engine {
Owned(Box<Memory<'static>>),
Mapped(OverlayMap),
}
impl Engine {
fn read<R>(&self, f: impl for<'a> FnOnce(&Memory<'a>) -> R) -> R {
match self {
Engine::Owned(mem) => f(mem),
Engine::Mapped(map) => f(map.borrow_dependent()),
}
}
fn with<R>(
&mut self,
store: &mut FileStorage,
f: impl for<'a> FnOnce(&mut Memory<'a>, &mut FileStorage) -> R,
) -> R {
match self {
Engine::Owned(mem) => f(mem, store),
Engine::Mapped(map) => map.with_dependent_mut(|_owner, mem| f(mem, store)),
}
}
}
fn open_engine(store: &mut FileStorage, cfg: &Config) -> Result<(Engine, OpenReport), HostError> {
let journal = store.read_journal()?;
let Some(genp) = store.current_snapshot_path()? else {
let (mem, report) = Memory::from_bytes(None, &journal, cfg.clone())?;
return Ok((Engine::Owned(Box::new(mem)), report));
};
let file = File::open(&genp).map_err(|e| HostError::io(&genp, e))?;
let map = unsafe { Mmap::map(&file) }.map_err(|e| HostError::io(&genp, e))?;
drop(file);
let mut report = None;
let mapped = OverlayMap::try_new(map, |m| {
let (mem, rep) = Memory::from_bytes_overlay(&m[..], &journal, cfg.clone())?;
report = Some(rep);
Ok::<_, Error>(mem)
})?;
Ok((Engine::Mapped(mapped), report.unwrap_or_default()))
}
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct FactSnapshot {
pub record: FactRecord,
pub text: String,
pub metadata: BTreeMap<String, String>,
}
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ExportedFact {
pub text: String,
pub entity: Option<String>,
pub tags: Vec<String>,
pub metadata: BTreeMap<String, String>,
pub recorded_at: u64,
pub valid_from: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct RecoverReport {
pub kept: usize,
pub dropped_text: usize,
pub dropped_vector: usize,
pub dropped_metadata: usize,
}
pub(crate) fn metadata_map(mem: &Memory, id: plugmem_core::FactId) -> BTreeMap<String, String> {
let mut pairs = Vec::new();
mem.metadata_of(id, &mut pairs);
pairs
.into_iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect()
}
pub(crate) fn export_facts_each(mem: &Memory, mut f: impl FnMut(ExportedFact)) {
use plugmem_core::{EntityId, FactId, VALID_TO_OPEN};
let next = mem.stats().next_fact;
let mut terms = Vec::new();
for i in 0..next {
let id = FactId(i);
let Some(view) = mem.get(id) else {
continue; };
if view.record.valid_to != VALID_TO_OPEN {
continue; }
let entity = (view.record.entity != EntityId::NONE)
.then(|| mem.entity_name(view.record.entity))
.flatten()
.map(str::to_string);
terms.clear();
mem.tags_of(id, &mut terms);
let tags = terms.iter().map(|t| mem.term(*t).to_string()).collect();
f(ExportedFact {
text: view.text.to_string(),
entity,
tags,
metadata: metadata_map(mem, id),
recorded_at: view.record.recorded_at,
valid_from: view.record.valid_from,
});
}
}
pub(crate) fn export_facts(mem: &Memory) -> Vec<ExportedFact> {
let mut out = Vec::new();
export_facts_each(mem, |e| out.push(e));
out
}
pub struct DatabaseBuilder {
cfg: Config,
fsync: FsyncPolicy,
snapshot_every_ops: u64,
snapshot_journal_bytes: u64,
maintain_every_forgets: Option<u64>,
embedder: Option<Box<dyn Embedder>>,
}
impl DatabaseBuilder {
pub fn fsync(mut self, policy: FsyncPolicy) -> Self {
self.fsync = policy;
self
}
pub fn snapshot_every_ops(mut self, ops: u64) -> Self {
self.snapshot_every_ops = ops;
self
}
pub fn snapshot_journal_bytes(mut self, bytes: u64) -> Self {
self.snapshot_journal_bytes = bytes;
self
}
pub fn maintain_every_forgets(mut self, forgets: u64) -> Self {
self.maintain_every_forgets = Some(forgets);
self
}
pub fn embedder(mut self, embedder: Box<dyn Embedder>) -> Self {
self.embedder = Some(embedder);
self
}
pub fn open(self, path: impl Into<PathBuf>) -> Result<(Database, OpenReport), HostError> {
if let Some(embedder) = &self.embedder {
let dim = embedder.dim();
if dim != 0 && dim != self.cfg.dim {
return Err(HostError::Engine(Error::ConfigMismatch(
"embedder dimension must equal Config::dim",
)));
}
}
let mut store = FileStorage::open(path, self.fsync)?;
let (engine, report) = open_engine(&mut store, &self.cfg)?;
let db = Database {
inner: Arc::new(Inner {
state: StateLock::new(State {
engine,
store,
ops: 0,
forgets: 0,
}),
embedder: self.embedder.map(Mutex::new),
cfg: self.cfg,
snapshot_every_ops: self.snapshot_every_ops,
snapshot_journal_bytes: self.snapshot_journal_bytes,
maintain_every_forgets: self.maintain_every_forgets,
}),
};
Ok((db, report))
}
}
#[cfg(not(feature = "counters"))]
type StateLock = RwLock<State>;
#[cfg(feature = "counters")]
type StateLock = Mutex<State>;
struct Inner {
state: StateLock,
embedder: Option<Mutex<Box<dyn Embedder>>>,
cfg: Config,
snapshot_every_ops: u64,
snapshot_journal_bytes: u64,
maintain_every_forgets: Option<u64>,
}
struct State {
engine: Engine,
store: FileStorage,
ops: u64,
forgets: u64,
}
#[derive(Clone)]
pub struct Database {
inner: Arc<Inner>,
}
impl Database {
pub fn open(path: impl Into<PathBuf>, cfg: Config) -> Result<(Self, OpenReport), HostError> {
Self::builder(cfg).open(path)
}
pub fn open_readonly(
path: impl Into<PathBuf>,
cfg: Config,
) -> Result<ReadOnlyDatabase, HostError> {
ReadOnlyDatabase::open(path, cfg)
}
pub fn builder(cfg: Config) -> DatabaseBuilder {
DatabaseBuilder {
cfg,
fsync: FsyncPolicy::default(),
snapshot_every_ops: 1024,
snapshot_journal_bytes: 4 * 1024 * 1024,
maintain_every_forgets: None,
embedder: None,
}
}
#[cfg(not(feature = "counters"))]
fn read(&self) -> RwLockReadGuard<'_, State> {
self.inner.state.read().unwrap_or_else(|e| e.into_inner())
}
#[cfg(not(feature = "counters"))]
fn write(&self) -> RwLockWriteGuard<'_, State> {
self.inner.state.write().unwrap_or_else(|e| e.into_inner())
}
#[cfg(feature = "counters")]
fn read(&self) -> MutexGuard<'_, State> {
self.inner.state.lock().unwrap_or_else(|e| e.into_inner())
}
#[cfg(feature = "counters")]
fn write(&self) -> MutexGuard<'_, State> {
self.inner.state.lock().unwrap_or_else(|e| e.into_inner())
}
fn embed_one(&self, text: &str) -> Result<Option<Vec<f32>>, HostError> {
let Some(embedder) = &self.inner.embedder else {
return Ok(None);
};
let mut embedder = embedder.lock().unwrap_or_else(|e| e.into_inner());
if embedder.dim() == 0 {
return Ok(None);
}
let mut vs = embedder.embed(&[text])?;
Ok(Some(vs.remove(0)))
}
fn embed_many(&self, texts: &[&str]) -> Result<Option<Vec<Vec<f32>>>, HostError> {
let Some(embedder) = &self.inner.embedder else {
return Ok(None);
};
let mut embedder = embedder.lock().unwrap_or_else(|e| e.into_inner());
if embedder.dim() == 0 {
return Ok(None);
}
if texts.is_empty() {
return Ok(Some(Vec::new()));
}
Ok(Some(embedder.embed(texts)?))
}
fn resnapshot(&self, st: &mut State, now: u64) -> Result<(), HostError> {
{
let State { engine, store, .. } = &mut *st;
store.stage_snapshot(|sink| {
engine
.read(|mem| mem.write_snapshot_to(now, &mut *sink))
.map_err(HostError::from)
})?;
}
st.engine = Engine::Owned(Box::new(Memory::new(self.inner.cfg.clone())?));
let write = st
.store
.commit_snapshot()
.and_then(|()| st.store.clear_journal());
let (engine, _) = open_engine(&mut st.store, &self.inner.cfg)?;
st.engine = engine;
write
}
fn after_mutation(&self, st: &mut State, now: u64) -> Result<(), HostError> {
st.ops += 1;
if let Some(threshold) = self.inner.maintain_every_forgets
&& st.forgets >= threshold
{
let State { engine, store, .. } = &mut *st;
engine.with(store, |mem, store| mem.maintain(store, now))?;
st.forgets = 0;
}
let by_ops = self.inner.snapshot_every_ops > 0 && st.ops >= self.inner.snapshot_every_ops;
let by_bytes = self.inner.snapshot_journal_bytes > 0
&& st.store.journal_bytes() >= self.inner.snapshot_journal_bytes;
if by_ops || by_bytes {
self.resnapshot(st, now)?;
st.ops = 0;
}
Ok(())
}
pub fn remember(&self, input: RememberInput<'_>) -> Result<RememberOutcome, HostError> {
let embedded = match input.vector {
Some(_) => None,
None => self.embed_one(input.text)?,
};
let input = RememberInput {
vector: embedded.as_deref().or(input.vector),
..input
};
let mut st = self.write();
let State { engine, store, .. } = &mut *st;
let out = engine.with(store, |mem, store| mem.remember(store, input))?;
self.after_mutation(&mut st, input.now)?;
Ok(out)
}
pub fn remember_many(
&self,
inputs: Vec<RememberInput<'_>>,
) -> Result<Vec<RememberOutcome>, HostError> {
if inputs.is_empty() {
return Ok(Vec::new());
}
let to_embed: Vec<&str> = inputs
.iter()
.filter(|i| i.vector.is_none())
.map(|i| i.text)
.collect();
let embedded = if to_embed.is_empty() {
None
} else {
self.embed_many(&to_embed)?
};
let mut st = self.write();
st.store.set_batch(true);
let mut out = Vec::with_capacity(inputs.len());
let mut cursor = 0usize; let mut latest = 0u64;
let mut failed = None;
for input in inputs {
latest = latest.max(input.now);
let vector = if input.vector.is_some() {
input.vector
} else if let Some(embedded) = &embedded {
let v = embedded[cursor].as_slice();
cursor += 1;
Some(v)
} else {
None };
let input = RememberInput { vector, ..input };
let State { engine, store, .. } = &mut *st;
match engine.with(store, |mem, store| mem.remember(store, input)) {
Ok(o) => out.push(o),
Err(e) => {
failed = Some(HostError::from(e));
break;
}
}
}
st.store.set_batch(false);
st.store.sync_journal()?;
if let Some(e) = failed {
return Err(e);
}
self.after_mutation(&mut st, latest)?;
Ok(out)
}
pub fn recall(&self, q: RecallQuery<'_>) -> Result<RecallResult, HostError> {
let embedded = match (q.vector, q.text) {
(None, Some(text)) => self.embed_one(text)?,
_ => None,
};
let q = RecallQuery {
vector: embedded.as_deref().or(q.vector),
..q
};
let st = self.read();
RECALL_SCRATCH.with(|scratch| {
let mut scratch = scratch.borrow_mut();
let mut out = RecallResult::default();
st.engine
.read(|mem| mem.recall_into(q, &mut scratch, &mut out))?;
Ok(out)
})
}
pub fn revise(
&self,
target: plugmem_core::FactId,
input: RememberInput<'_>,
) -> Result<RememberOutcome, HostError> {
let embedded = match input.vector {
Some(_) => None,
None => self.embed_one(input.text)?,
};
let input = RememberInput {
vector: embedded.as_deref().or(input.vector),
..input
};
let mut st = self.write();
let State { engine, store, .. } = &mut *st;
let out = engine.with(store, |mem, store| mem.revise(store, target, input))?;
self.after_mutation(&mut st, input.now)?;
Ok(out)
}
pub fn forget(&self, now: u64, id: plugmem_core::FactId) -> Result<bool, HostError> {
let mut st = self.write();
let State { engine, store, .. } = &mut *st;
let fresh = engine.with(store, |mem, store| mem.forget(store, now, id))?;
st.forgets += 1;
self.after_mutation(&mut st, now)?;
Ok(fresh)
}
pub fn link(&self, input: LinkInput<'_>) -> Result<(), HostError> {
let mut st = self.write();
let State { engine, store, .. } = &mut *st;
engine.with(store, |mem, store| mem.link(store, input))?;
self.after_mutation(&mut st, input.now)?;
Ok(())
}
pub fn get(&self, id: plugmem_core::FactId) -> Option<FactSnapshot> {
self.read().engine.read(|mem| {
mem.get(id).map(|v| FactSnapshot {
record: v.record,
text: v.text.to_string(),
metadata: metadata_map(mem, id),
})
})
}
pub fn stats(&self) -> Stats {
self.read().engine.read(|mem| mem.stats())
}
pub fn export(&self) -> Vec<ExportedFact> {
self.read().engine.read(export_facts)
}
pub fn export_each(&self, f: impl FnMut(ExportedFact)) {
self.read().engine.read(|mem| export_facts_each(mem, f));
}
pub fn maintain(&self, now: u64) -> Result<MaintainReport, HostError> {
let mut st = self.write();
let snap_len = |store: &FileStorage| -> usize {
store
.current_snapshot_path()
.ok()
.flatten()
.and_then(|p| std::fs::metadata(&p).ok())
.map(|m| m.len() as usize)
.unwrap_or(0)
};
let bytes_before = snap_len(&st.store);
let text_tmp = tmp_sibling(st.store.path(), "mtext");
let vec_tmp = tmp_sibling(st.store.path(), "mvec");
let mut purged = 0usize;
{
let State { engine, store, .. } = &mut *st;
store.stage_snapshot(|sink| {
engine.read(|mem| {
let mut text_scratch = FileScratch::create(&text_tmp)?;
let mut vec_scratch = FileScratch::create(&vec_tmp)?;
purged = mem
.snapshot_disk_first(now, &mut text_scratch, &mut vec_scratch, &mut *sink)
.map_err(HostError::from)?;
Ok(())
})
})?;
}
st.engine = Engine::Owned(Box::new(Memory::new(self.inner.cfg.clone())?));
st.store
.commit_snapshot()
.and_then(|()| st.store.clear_journal())?;
let (engine, _) = open_engine(&mut st.store, &self.inner.cfg)?;
st.engine = engine;
st.forgets = 0;
st.ops = 0;
let bytes_after = snap_len(&st.store);
Ok(MaintainReport {
purged,
bytes_before,
bytes_after,
})
}
pub fn checkpoint(&self, now: u64) -> Result<(), HostError> {
let mut st = self.write();
self.resnapshot(&mut st, now)?;
st.ops = 0;
Ok(())
}
pub fn verify(&self) -> Result<(), HostError> {
Ok(self.read().engine.read(|mem| mem.verify())?)
}
pub fn recover(
src: impl AsRef<Path>,
dst: impl AsRef<Path>,
cfg: Config,
now: u64,
) -> Result<RecoverReport, HostError> {
let src = src.as_ref();
let dst = dst.as_ref();
let mut src_store = FileStorage::open(src, FsyncPolicy::OnSnapshot)?;
let src_base = src_store.path().to_path_buf();
let same = dst == src_base
|| matches!(
(std::fs::canonicalize(dst), std::fs::canonicalize(&src_base)),
(Ok(a), Ok(b)) if a == b
);
if same {
return Err(HostError::Engine(Error::Invalid(
"recover destination must differ from the source",
)));
}
let journal = src_store.read_journal()?;
let Some(genp) = src_store.current_snapshot_path()? else {
return Err(HostError::Engine(Error::Corrupt(
"source database has no published snapshot to recover",
)));
};
let file = File::open(&genp).map_err(|e| HostError::io(&genp, e))?;
let map = unsafe { Mmap::map(&file) }.map_err(|e| HostError::io(&genp, e))?;
drop(file);
let (mut mem, _report) = Memory::from_bytes_overlay(&map[..], &journal, cfg.clone())?;
let mut scratch = MemStorage::new();
let mut dropped_text = 0usize;
let mut dropped_vector = 0usize;
let mut dropped_metadata = 0usize;
for (id, fault) in mem.faulty_facts() {
mem.forget(&mut scratch, now, id)?;
match fault {
FactFault::Text => dropped_text += 1,
FactFault::Vector => dropped_vector += 1,
FactFault::Metadata => dropped_metadata += 1,
}
}
let mut dst_store = FileStorage::open(dst, FsyncPolicy::OnSnapshot)?;
let text_tmp = tmp_sibling(dst_store.path(), "rectext");
let vec_tmp = tmp_sibling(dst_store.path(), "recvec");
let mut purged = 0usize;
dst_store.stage_snapshot(|sink| {
let mut text_scratch = FileScratch::create(&text_tmp)?;
let mut vec_scratch = FileScratch::create(&vec_tmp)?;
purged = mem
.snapshot_disk_first(now, &mut text_scratch, &mut vec_scratch, &mut *sink)
.map_err(HostError::from)?;
Ok(())
})?;
dst_store.commit_snapshot()?;
let kept = mem.stats().facts.saturating_sub(purged);
Ok(RecoverReport {
kept,
dropped_text,
dropped_vector,
dropped_metadata,
})
}
}
fn tmp_sibling(base: &Path, tag: &str) -> PathBuf {
let mut p = base.as_os_str().to_os_string();
p.push(".");
p.push(tag);
p.push(".tmp");
PathBuf::from(p)
}
impl std::fmt::Debug for Database {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let stats = self.stats();
f.debug_struct("Database")
.field("facts", &stats.facts)
.field("entities", &stats.entities)
.finish()
}
}