use std::collections::HashMap;
use std::sync::Arc;
use rusqlite::Connection;
use crate::embed::Embedder;
use crate::scorer::score_with_lexical;
use crate::{Memory, Query, ScoredMemory, Weights};
fn bm25_similarities(memories: &[Memory], query: &str) -> HashMap<i64, f32> {
let mut out = HashMap::new();
let match_expr = fts_match_expr(query);
if match_expr.is_empty() {
return out;
}
let conn = match Connection::open_in_memory() {
Ok(c) => c,
Err(_) => return out,
};
if conn
.execute("CREATE VIRTUAL TABLE mem_fts USING fts5(text)", [])
.is_err()
{
return out;
}
{
let mut ins = match conn.prepare("INSERT INTO mem_fts(rowid, text) VALUES (?1, ?2)") {
Ok(s) => s,
Err(_) => return out,
};
for m in memories {
let _ = ins.execute(rusqlite::params![m.id, m.text]);
}
}
let mut stmt = match conn
.prepare("SELECT rowid, bm25(mem_fts) FROM mem_fts WHERE mem_fts MATCH ?1")
{
Ok(s) => s,
Err(_) => return out,
};
let rows = stmt.query_map([&match_expr], |r| {
Ok((r.get::<_, i64>(0)?, r.get::<_, f64>(1)?))
});
if let Ok(rows) = rows {
for (id, bm25) in rows.flatten() {
let x = (-bm25).max(0.0) as f32; out.insert(id, x / (1.0 + x));
}
}
out
}
fn fts_match_expr(query: &str) -> String {
let mut seen = std::collections::BTreeSet::new();
query
.split(|c: char| !c.is_alphanumeric())
.map(|w| w.to_lowercase())
.filter(|w| w.len() >= 2 && seen.insert(w.clone()))
.map(|t| format!("\"{t}\""))
.collect::<Vec<_>>()
.join(" OR ")
}
pub trait MemoryEngine {
fn name(&self) -> &str;
fn retrieve(&self, query: &Query, k: usize) -> Vec<ScoredMemory>;
fn explain(&self, id: i64, query: &Query) -> Option<ScoredMemory>;
}
pub struct BuiltinEngine {
pub memories: Vec<Memory>,
pub weights: Weights,
embedder: Option<Arc<dyn Embedder>>,
}
impl BuiltinEngine {
pub fn new(memories: Vec<Memory>) -> Self {
Self {
memories,
weights: Weights::default(),
embedder: None,
}
}
pub fn with_weights(mut self, weights: Weights) -> Self {
self.weights = weights;
self
}
pub fn with_embedder(mut self, embedder: Arc<dyn Embedder>) -> anyhow::Result<Self> {
for m in &mut self.memories {
if m.embedding.is_none() {
m.embedding = Some(embedder.embed(&m.text)?);
}
}
self.embedder = Some(embedder);
Ok(self)
}
fn query_embedding(&self, query: &Query) -> Option<Vec<f32>> {
self.embedder.as_ref().and_then(|e| e.embed(&query.text).ok())
}
}
impl MemoryEngine for BuiltinEngine {
fn name(&self) -> &str {
if self.embedder.is_some() {
"builtin+embeddings"
} else {
"builtin"
}
}
fn retrieve(&self, query: &Query, k: usize) -> Vec<ScoredMemory> {
let qe = self.query_embedding(query);
let sims = if qe.is_none() {
bm25_similarities(&self.memories, &query.text)
} else {
HashMap::new()
};
let mut scored: Vec<ScoredMemory> = self
.memories
.iter()
.map(|m| {
score_with_lexical(m, query, &self.weights, qe.as_deref(), sims.get(&m.id).copied())
})
.collect();
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
scored.truncate(k);
scored
}
fn explain(&self, id: i64, query: &Query) -> Option<ScoredMemory> {
let qe = self.query_embedding(query);
let sims = if qe.is_none() {
bm25_similarities(&self.memories, &query.text)
} else {
HashMap::new()
};
self.memories
.iter()
.find(|m| m.id == id)
.map(|m| {
score_with_lexical(m, query, &self.weights, qe.as_deref(), sims.get(&m.id).copied())
})
}
}
#[cfg(feature = "mempalace")]
pub struct MemPalaceEngine {
pub command: String,
pub args: Vec<String>,
pub tool: String,
}
#[cfg(feature = "mempalace")]
impl Default for MemPalaceEngine {
fn default() -> Self {
Self {
command: "mempalace-mcp".into(),
args: Vec::new(),
tool: "search".into(),
}
}
}
#[cfg(feature = "mempalace")]
impl MemPalaceEngine {
pub fn new(command: impl Into<String>, args: Vec<String>) -> Self {
Self {
command: command.into(),
args,
tool: "search".into(),
}
}
pub fn with_tool(mut self, tool: impl Into<String>) -> Self {
self.tool = tool.into();
self
}
pub fn try_retrieve(&self, query: &Query, k: usize) -> anyhow::Result<Vec<ScoredMemory>> {
let mut client = crate::mcp::McpClient::spawn(&self.command, &self.args)?;
let tools = client.list_tools()?;
if let Some(names) = tools.get("tools").and_then(serde_json::Value::as_array) {
if !names
.iter()
.any(|t| t.get("name").and_then(serde_json::Value::as_str) == Some(&self.tool))
{
anyhow::bail!(
"`{}` does not expose a `{}` tool",
self.command,
self.tool
);
}
}
let text = client.call_tool(
&self.tool,
serde_json::json!({"query": query.text, "limit": k}),
)?;
let mut out = map_hits(&text, query)?;
out.truncate(k);
Ok(out)
}
}
#[cfg(feature = "mempalace")]
fn map_hits(payload: &str, query: &Query) -> anyhow::Result<Vec<ScoredMemory>> {
use anyhow::Context;
use serde_json::Value;
let parsed: Value = serde_json::from_str(payload)
.context("mempalace search result was not JSON")?;
let hits = match &parsed {
Value::Array(a) => a.clone(),
Value::Object(o) => o
.get("results")
.and_then(Value::as_array)
.cloned()
.ok_or_else(|| anyhow::anyhow!("mempalace search result has no `results` array"))?,
_ => anyhow::bail!("unexpected mempalace search result shape"),
};
let ts = |h: &Value, key: &str| {
h.get(key)
.and_then(Value::as_str)
.and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok())
.map(|d| d.with_timezone(&chrono::Utc))
.unwrap_or(query.now)
};
Ok(hits
.iter()
.enumerate()
.map(|(i, h)| {
let score = h.get("score").and_then(Value::as_f64).unwrap_or(0.0) as f32;
let score = score.clamp(0.0, 1.0);
let memory = Memory {
id: h.get("id").and_then(Value::as_i64).unwrap_or(i as i64),
text: h
.get("text")
.and_then(Value::as_str)
.unwrap_or_default()
.to_string(),
created_at: ts(h, "created_at"),
last_used: ts(h, "last_used"),
mentions: h.get("mentions").and_then(Value::as_u64).unwrap_or(0) as u32,
importance: h.get("importance").and_then(Value::as_f64).unwrap_or(0.0) as f32,
tags: h
.get("tags")
.and_then(Value::as_array)
.map(|t| {
t.iter()
.filter_map(Value::as_str)
.map(str::to_string)
.collect()
})
.unwrap_or_default(),
embedding: None,
};
ScoredMemory {
memory,
score,
signals: vec![crate::Signal {
name: "similarity".into(),
weight: 1.0,
score,
applicable: true,
detail: format!("mempalace semantic score {score:.2}"),
}],
}
})
.collect())
}
#[cfg(feature = "mempalace")]
impl MemoryEngine for MemPalaceEngine {
fn name(&self) -> &str {
"mempalace"
}
fn retrieve(&self, query: &Query, k: usize) -> Vec<ScoredMemory> {
match self.try_retrieve(query, k) {
Ok(hits) => hits,
Err(e) => {
eprintln!("[mw-memory] mempalace retrieval failed: {e:#}");
Vec::new()
}
}
}
fn explain(&self, id: i64, query: &Query) -> Option<ScoredMemory> {
self.try_retrieve(query, 50)
.ok()?
.into_iter()
.find(|s| s.memory.id == id)
}
}
#[cfg(feature = "mempalace")]
pub struct Drawer {
pub wing: String,
pub room: String,
pub content: String,
}
#[cfg(feature = "mempalace")]
#[derive(Debug, Default, PartialEq, Eq)]
pub struct CheckpointOutcome {
pub added: usize,
pub duplicates: usize,
pub errors: usize,
}
#[cfg(feature = "mempalace")]
pub fn checkpoint(
command: &str,
args: &[String],
tool: &str,
items: &[Drawer],
) -> anyhow::Result<CheckpointOutcome> {
use serde_json::{json, Value};
let mut client = crate::mcp::McpClient::spawn(command, args)?;
let payload = json!({
"items": items
.iter()
.map(|d| json!({"wing": d.wing, "room": d.room, "content": d.content}))
.collect::<Vec<_>>(),
"added_by": "memorywhale",
});
let text = client.call_tool(tool, payload)?;
let v: Value = serde_json::from_str(&text).unwrap_or(Value::Null);
let count = |k: &str| v.get(k).and_then(Value::as_array).map(Vec::len).unwrap_or(0);
Ok(CheckpointOutcome {
added: count("added"),
duplicates: count("duplicates"),
errors: count("errors"),
})
}
#[cfg(feature = "mempalace")]
pub enum SyncOp {
Add {
wing: String,
room: String,
content: String,
added_by: String,
},
Delete {
drawer_id: String,
},
}
#[cfg(feature = "mempalace")]
pub enum SyncResult {
Added { drawer_id: String },
Deleted,
}
#[cfg(feature = "mempalace")]
pub fn sync_ops(
command: &str,
args: &[String],
add_tool: &str,
delete_tool: &str,
ops: &[SyncOp],
) -> anyhow::Result<Vec<SyncResult>> {
use anyhow::Context;
use serde_json::{json, Value};
let mut client = crate::mcp::McpClient::spawn(command, args)?;
let mut out = Vec::with_capacity(ops.len());
for op in ops {
match op {
SyncOp::Add { wing, room, content, added_by } => {
let text = client.call_tool(
add_tool,
json!({"wing": wing, "room": room, "content": content, "added_by": added_by}),
)?;
let v: Value =
serde_json::from_str(&text).context("add_drawer result was not JSON")?;
let drawer_id = v
.get("drawer_id")
.and_then(Value::as_str)
.ok_or_else(|| anyhow::anyhow!("add_drawer returned no drawer_id"))?
.to_string();
out.push(SyncResult::Added { drawer_id });
}
SyncOp::Delete { drawer_id } => {
client.call_tool(delete_tool, json!({"drawer_id": drawer_id}))?;
out.push(SyncResult::Deleted);
}
}
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, TimeZone, Utc};
fn now() -> chrono::DateTime<Utc> {
Utc.with_ymd_and_hms(2026, 6, 27, 12, 0, 0).unwrap()
}
fn sample() -> Vec<Memory> {
let n = now();
vec![
Memory { id: 143, text: "I use Rust for systems software.".into(), created_at: n - Duration::days(20), last_used: n, mentions: 27, importance: 0.98, tags: vec!["rust".into()], embedding: None },
Memory { id: 7, text: "I ate pizza.".into(), created_at: n - Duration::days(40), last_used: n - Duration::days(40), mentions: 1, importance: 0.01, tags: vec![], embedding: None },
Memory { id: 22, text: "Use Tokio for async runtime.".into(), created_at: n - Duration::days(5), last_used: n - Duration::days(3), mentions: 6, importance: 0.6, tags: vec!["rust".into(), "tokio".into()], embedding: None },
]
}
#[test]
fn retrieve_ranks_and_truncates() {
let eng = BuiltinEngine::new(sample());
let q = Query::new("which language for systems programming?", now());
let top = eng.retrieve(&q, 2);
assert_eq!(top.len(), 2);
assert_eq!(top[0].memory.id, 143); }
#[test]
fn bm25_similarity_signal_fires() {
let eng = BuiltinEngine::new(sample());
let q = Query::new("pizza", now());
let matching = eng.explain(7, &q).unwrap(); let sim_m = matching.signals.iter().find(|s| s.name == "similarity").unwrap();
assert!(sim_m.detail.contains("BM25"), "should use BM25: {}", sim_m.detail);
assert!(sim_m.score > 0.0, "matching memory sim should fire: {}", sim_m.score);
let non = eng.explain(143, &q).unwrap(); let sim_n = non.signals.iter().find(|s| s.name == "similarity").unwrap();
assert_eq!(sim_n.score, 0.0, "non-matching sim should be 0: {}", sim_n.score);
assert!(sim_m.score > sim_n.score);
}
#[test]
fn explain_returns_breakdown() {
let eng = BuiltinEngine::new(sample());
let q = Query::new("rust", now());
let e = eng.explain(143, &q).unwrap();
assert!(e.explain().contains("memory explain 143"));
assert!(eng.explain(9999, &q).is_none());
}
#[cfg(feature = "mempalace")]
#[test]
fn mempalace_maps_hits_to_signals() {
let payload = include_str!("../tests/fixtures/mempalace_search.json");
let hits = map_hits(payload, &Query::new("rust", now())).unwrap();
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].memory.id, 143);
assert_eq!(hits[0].memory.tags, vec!["rust".to_string()]);
assert_eq!(hits[0].percent(), 87);
let sig = &hits[0].signals[0];
assert_eq!(sig.name, "similarity");
assert_eq!(sig.detail, "mempalace semantic score 0.87");
assert_eq!(hits[1].memory.created_at, now());
}
#[cfg(feature = "mempalace")]
#[test]
fn mempalace_rejects_garbage() {
assert!(map_hits("not json", &Query::new("x", now())).is_err());
assert!(map_hits("{\"oops\": 1}", &Query::new("x", now())).is_err());
}
#[cfg(all(feature = "mempalace", unix))]
#[test]
fn mempalace_talks_to_a_fake_server() {
let script = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/fake-mempalace-mcp.sh");
let eng = MemPalaceEngine::new("sh", vec![script.to_string()]);
let hits = eng.try_retrieve(&Query::new("rust", now()), 10).unwrap();
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].memory.id, 143);
assert!(hits[0].reasons()[0].starts_with("mempalace semantic score"));
}
#[cfg(all(feature = "mempalace", unix))]
#[test]
fn checkpoint_pushes_and_parses_the_summary() {
let script =
concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/fake-mempalace-checkpoint.sh");
let items = vec![
Drawer { wing: "memorywhale".into(), room: "command".into(), content: "cargo build failed".into() },
Drawer { wing: "memorywhale".into(), room: "note".into(), content: "the fix was X".into() },
];
let out = checkpoint("sh", &[script.to_string()], "mempalace_checkpoint", &items).unwrap();
assert_eq!(out, CheckpointOutcome { added: 2, duplicates: 1, errors: 0 });
}
#[cfg(all(feature = "mempalace", unix))]
#[test]
fn sync_ops_adds_and_deletes_over_one_session() {
let script =
concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/fake-mempalace-sync.sh");
let ops = vec![
SyncOp::Add {
wing: "memorywhale".into(),
room: "note".into(),
content: "first".into(),
added_by: "you".into(),
},
SyncOp::Delete { drawer_id: "old-1".into() },
SyncOp::Add {
wing: "memorywhale".into(),
room: "note".into(),
content: "second".into(),
added_by: "memorywhale".into(),
},
];
let out = sync_ops(
"sh",
&[script.to_string()],
"mempalace_add_drawer",
"mempalace_delete_drawer",
&ops,
)
.unwrap();
assert_eq!(out.len(), 3);
let ids: Vec<&str> = out
.iter()
.filter_map(|r| match r {
SyncResult::Added { drawer_id } => Some(drawer_id.as_str()),
SyncResult::Deleted => None,
})
.collect();
assert_eq!(ids, vec!["drawer-1", "drawer-2"]);
assert!(matches!(out[1], SyncResult::Deleted));
}
#[cfg(feature = "mempalace")]
#[test]
fn missing_server_is_a_clear_error() {
let eng = MemPalaceEngine::new("mw-definitely-not-installed", vec![]);
let err = eng
.try_retrieve(&Query::new("x", now()), 3)
.unwrap_err()
.to_string();
assert!(err.contains("failed to start MCP server"), "{err}");
}
}