use std::path::{Path, PathBuf};
use crate::{Config, Database, DatabaseBuilder, Embedder, HostError, OpenAiCompatEmbedder};
const ENV_CONFIG: &str = "PLUGMEM_CONFIG";
const ENV_EMBEDDER: &str = "PLUGMEM_EMBEDDER";
const CONFIG_DIR: &str = "plugmem";
const CONFIG_FILE: &str = "config.toml";
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum SettingsError {
#[error("{0}")]
Config(String),
}
impl SettingsError {
fn config(msg: impl Into<String>) -> Self {
SettingsError::Config(msg.into())
}
}
pub struct Settings {
pub config: Config,
pub embedder: Option<Box<dyn Embedder>>,
pub snapshot_every_ops: Option<u64>,
pub snapshot_journal_bytes: Option<u64>,
pub maintain_every_forgets: Option<u64>,
}
impl Settings {
pub fn load(flag: Option<&Path>) -> Result<Settings, SettingsError> {
let table = read_config(flag)?;
Settings::from_table(table.as_ref())
}
pub fn from_table(table: Option<&toml::Table>) -> Result<Settings, SettingsError> {
let mut config = Config::default();
let mut embedder = EmbedderCfg::default();
let mut snapshot_every_ops = None;
let mut snapshot_journal_bytes = None;
let mut maintain_every_forgets = None;
if let Some(table) = table {
if let Some(t) = table.get("engine").and_then(toml::Value::as_table) {
apply_engine(&mut config, t)?;
}
if let Some(t) = table.get("embedder").and_then(toml::Value::as_table) {
embedder.merge(t);
}
if let Some(t) = table.get("maintenance").and_then(toml::Value::as_table) {
snapshot_every_ops = table_u64(t, "snapshot_every_ops");
snapshot_journal_bytes = table_u64(t, "snapshot_journal_bytes");
maintain_every_forgets = table_u64(t, "maintain_every_forgets");
}
}
if let Some(kind) = std::env::var_os(ENV_EMBEDDER) {
embedder.kind = Some(kind.to_string_lossy().into_owned());
}
let embedder = embedder.build(config.dim)?;
Ok(Settings {
config,
embedder,
snapshot_every_ops,
snapshot_journal_bytes,
maintain_every_forgets,
})
}
pub fn open(self, path: &Path) -> Result<Database, HostError> {
let mut b: DatabaseBuilder = Database::builder(self.config);
if let Some(v) = self.snapshot_every_ops {
b = b.snapshot_every_ops(v);
}
if let Some(v) = self.snapshot_journal_bytes {
b = b.snapshot_journal_bytes(v);
}
if let Some(v) = self.maintain_every_forgets {
b = b.maintain_every_forgets(v);
}
if let Some(e) = self.embedder {
b = b.embedder(e);
}
Ok(b.open(path)?.0)
}
}
pub fn read_config(flag: Option<&Path>) -> Result<Option<toml::Table>, SettingsError> {
let text = match read_config_text(flag)? {
Some(t) => t,
None => return Ok(None),
};
let table: toml::Table = text
.parse()
.map_err(|e| SettingsError::config(format!("config.toml is not valid TOML: {e}")))?;
Ok(Some(table))
}
pub(crate) fn table_u64(t: &toml::Table, key: &str) -> Option<u64> {
t.get(key)
.and_then(toml::Value::as_integer)
.filter(|n| *n >= 0)
.map(|n| n as u64)
}
fn read_config_text(flag: Option<&Path>) -> Result<Option<String>, SettingsError> {
if let Some(p) = flag {
return std::fs::read_to_string(p)
.map(Some)
.map_err(|e| SettingsError::config(format!("reading config {}: {e}", p.display())));
}
let candidate = std::env::var_os(ENV_CONFIG)
.map(PathBuf::from)
.or_else(default_config_path);
match candidate {
Some(p) if p.exists() => std::fs::read_to_string(&p)
.map(Some)
.map_err(|e| SettingsError::config(format!("reading config {}: {e}", p.display()))),
_ => Ok(None),
}
}
fn default_config_path() -> Option<PathBuf> {
std::env::var_os("XDG_CONFIG_HOME")
.map(PathBuf::from)
.or_else(|| std::env::var_os("HOME").map(|h| PathBuf::from(h).join(".config")))
.map(|base| base.join(CONFIG_DIR).join(CONFIG_FILE))
}
fn apply_engine(cfg: &mut Config, t: &toml::Table) -> Result<(), SettingsError> {
let fields: [(&str, &mut usize); 9] = [
("dim", &mut cfg.dim),
("max_bytes", &mut cfg.max_bytes),
("max_text", &mut cfg.max_text),
("max_blob", &mut cfg.max_blob),
("shards_facts", &mut cfg.shards_facts),
("shards_entities", &mut cfg.shards_entities),
("shards_edges", &mut cfg.shards_edges),
("shards_temporal", &mut cfg.shards_temporal),
("shards_postings", &mut cfg.shards_postings),
];
for (key, slot) in fields {
if let Some(v) = t.get(key) {
let n = v.as_integer().filter(|n| *n >= 0).ok_or_else(|| {
SettingsError::config(format!("[engine].{key} must be a non-negative integer"))
})?;
*slot = n as usize;
}
}
Ok(())
}
#[derive(Default)]
struct EmbedderCfg {
kind: Option<String>,
url: Option<String>,
model: Option<String>,
api_key_env: Option<String>,
}
impl EmbedderCfg {
fn merge(&mut self, t: &toml::Table) {
let s = |t: &toml::Table, k: &str| t.get(k).and_then(toml::Value::as_str).map(String::from);
if let Some(v) = s(t, "kind") {
self.kind = Some(v);
}
if let Some(v) = s(t, "url") {
self.url = Some(v);
}
if let Some(v) = s(t, "model") {
self.model = Some(v);
}
if let Some(v) = s(t, "api_key_env") {
self.api_key_env = Some(v);
}
}
fn build(&self, dim: usize) -> Result<Option<Box<dyn Embedder>>, SettingsError> {
let kind = self.kind.as_deref().unwrap_or("none");
match kind {
"none" | "" => Ok(None),
"ollama" | "openai" | "openai-compat" | "lmstudio" | "vllm" | "llamacpp" => {
let url = self.url.clone().ok_or_else(|| {
SettingsError::config(format!("[embedder] kind \"{kind}\" needs a url"))
})?;
let model = self.model.clone().ok_or_else(|| {
SettingsError::config(format!("[embedder] kind \"{kind}\" needs a model"))
})?;
if dim == 0 {
return Err(SettingsError::config(
"[embedder] requires [engine].dim > 0 (the embedding size)",
));
}
let mut e = OpenAiCompatEmbedder::new(&url, &model, dim);
if let Some(env) = &self.api_key_env
&& let Some(key) = std::env::var_os(env)
{
e = e.with_api_key(key.to_string_lossy().into_owned());
}
Ok(Some(Box::new(e)))
}
other => Err(SettingsError::config(format!(
"unknown [embedder] kind: {other}"
))),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct TempDir(PathBuf);
impl TempDir {
fn new(tag: &str) -> Self {
let dir = std::env::temp_dir().join(format!(
"plugmem-settings-{tag}-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&dir).unwrap();
TempDir(dir)
}
}
impl Drop for TempDir {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
#[test]
fn engine_and_maintenance_parse() {
let text = "\
[engine]
dim = 384
shards_facts = 16
[maintenance]
snapshot_every_ops = 50
snapshot_journal_bytes = 8192
maintain_every_forgets = 3
";
let table: toml::Table = text.parse().unwrap();
let s = Settings::from_table(Some(&table)).unwrap();
assert_eq!(s.config.dim, 384);
assert_eq!(s.config.shards_facts, 16);
assert_eq!(s.snapshot_every_ops, Some(50));
assert_eq!(s.snapshot_journal_bytes, Some(8192));
assert_eq!(s.maintain_every_forgets, Some(3));
let bad: toml::Table = "[engine]\ndim = \"huge\"".parse().unwrap();
assert!(matches!(
Settings::from_table(Some(&bad)),
Err(SettingsError::Config(_))
));
}
#[test]
fn defaults_when_no_table() {
let s = Settings::from_table(None).unwrap();
assert_eq!(s.config.dim, Config::default().dim);
assert!(s.embedder.is_none());
assert_eq!(s.snapshot_every_ops, None);
}
#[test]
fn embedder_merge_reads_every_field() {
let text = "\
[embedder]
kind = \"ollama\"
url = \"http://localhost:11434/v1\"
model = \"nomic-embed-text\"
api_key_env = \"SOME_ENV\"
[engine]
dim = 8
";
let table: toml::Table = text.parse().unwrap();
let s = Settings::from_table(Some(&table)).unwrap();
assert!(s.embedder.is_some());
}
#[test]
fn settings_open_applies_maintenance_and_embedder() {
let tmp = TempDir::new("open");
let mut config = Config::default();
config.dim = 8;
let embedder = EmbedderCfg {
kind: Some("ollama".into()),
url: Some("http://127.0.0.1:0/v1".into()),
model: Some("m".into()),
api_key_env: None,
}
.build(8)
.unwrap();
assert!(embedder.is_some());
let settings = Settings {
config,
embedder,
snapshot_every_ops: Some(4),
snapshot_journal_bytes: Some(4096),
maintain_every_forgets: Some(2),
};
let db = settings.open(&tmp.0.join("m.plugmem")).unwrap();
assert_eq!(db.stats().facts, 0);
}
#[test]
fn embedder_build_rules() {
assert!(EmbedderCfg::default().build(0).unwrap().is_none());
let no_url = EmbedderCfg {
kind: Some("ollama".into()),
..Default::default()
};
assert!(matches!(no_url.build(384), Err(SettingsError::Config(_))));
let no_model = EmbedderCfg {
kind: Some("ollama".into()),
url: Some("http://x/v1".into()),
..Default::default()
};
assert!(matches!(no_model.build(384), Err(SettingsError::Config(_))));
let zero_dim = EmbedderCfg {
kind: Some("ollama".into()),
url: Some("http://x/v1".into()),
model: Some("m".into()),
api_key_env: None,
};
assert!(matches!(zero_dim.build(0), Err(SettingsError::Config(_))));
let ok = EmbedderCfg {
kind: Some("openai".into()),
url: Some("http://x/v1".into()),
model: Some("m".into()),
api_key_env: Some("PLUGMEM_TEST_KEY_UNSET".into()),
};
assert!(ok.build(384).unwrap().is_some());
let weird = EmbedderCfg {
kind: Some("weird".into()),
..Default::default()
};
assert!(matches!(weird.build(384), Err(SettingsError::Config(_))));
}
#[test]
fn load_reads_the_config_file() {
let tmp = TempDir::new("load");
let cfgfile = tmp.0.join("config.toml");
std::fs::write(
&cfgfile,
"[engine]\ndim = 512\n[embedder]\nkind = \"none\"\n[maintenance]\nsnapshot_every_ops = 64\n",
)
.unwrap();
let s = Settings::load(Some(&cfgfile)).unwrap();
assert_eq!(s.config.dim, 512);
assert!(s.embedder.is_none());
assert_eq!(s.snapshot_every_ops, Some(64));
assert!(matches!(
Settings::load(Some(&tmp.0.join("nope.toml"))),
Err(SettingsError::Config(_))
));
}
#[test]
fn read_config_none_and_batch_extra() {
let tmp = TempDir::new("extra");
let missing = tmp.0.join("absent.toml");
assert!(read_config(Some(&missing)).is_err());
let cfgfile = tmp.0.join("config.toml");
std::fs::write(&cfgfile, "[maintenance]\nbatch_size = 256\n").unwrap();
let table = read_config(Some(&cfgfile)).unwrap().unwrap();
let batch = table
.get("maintenance")
.and_then(toml::Value::as_table)
.and_then(|m| table_u64(m, "batch_size"));
assert_eq!(batch, Some(256));
}
}