use std::{collections::BTreeSet, path::Path, time::Duration};
use rmcp::schemars;
use rusqlite::{Connection, OptionalExtension, Result, ToSql, params};
use serde::Serialize;
#[derive(Debug, Clone, Serialize, schemars::JsonSchema)]
pub struct Memory {
pub key: String,
pub repo: Option<String>,
pub tags: Vec<String>,
pub content: String,
pub created_at: String,
pub updated_at: String,
}
#[derive(Debug)]
pub struct Database {
conn: Connection,
}
impl Database {
pub fn open(path: impl AsRef<Path>) -> Result<Self> {
let conn = Connection::open(path)?;
conn.busy_timeout(Duration::from_secs(5))?;
conn.pragma_update(None, "foreign_keys", "ON")?;
conn.pragma_update(None, "journal_mode", "WAL")?;
let db = Self { conn };
db.init()?;
Ok(db)
}
#[cfg(test)]
fn open_in_memory() -> Result<Self> {
let conn = Connection::open_in_memory()?;
let db = Self { conn };
db.init()?;
Ok(db)
}
pub fn init(&self) -> Result<()> {
self.conn.execute_batch(
"
CREATE TABLE IF NOT EXISTS memories (
key TEXT PRIMARY KEY NOT NULL,
repo TEXT,
content TEXT NOT NULL,
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now')),
updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now'))
);
CREATE TABLE IF NOT EXISTS memory_tags (
memory_key TEXT NOT NULL,
tag TEXT NOT NULL,
PRIMARY KEY (memory_key, tag),
FOREIGN KEY (memory_key) REFERENCES memories(key) ON DELETE CASCADE
);
",
)?;
self.conn.execute_batch(
"
CREATE INDEX IF NOT EXISTS idx_memories_repo
ON memories(repo, updated_at);
CREATE INDEX IF NOT EXISTS idx_memories_updated_at
ON memories(updated_at);
CREATE INDEX IF NOT EXISTS idx_memory_tags_tag
ON memory_tags(tag, memory_key);
",
)
}
pub fn upsert_memory(
&mut self,
key: &str,
repo: Option<&str>,
tags: &[String],
content: &str,
) -> Result<Memory> {
let repo = normalize_repo(repo);
let tags = normalize_tags(tags);
let tx = self.conn.transaction()?;
tx.execute("PRAGMA foreign_keys = ON", [])?;
tx.execute(
"
INSERT INTO memories (key, repo, content)
VALUES (?1, ?2, ?3)
ON CONFLICT(key) DO UPDATE SET
repo = excluded.repo,
content = excluded.content,
updated_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now')
",
params![key, repo, content],
)?;
tx.execute(
"DELETE FROM memory_tags WHERE memory_key = ?1",
params![key],
)?;
for tag in tags {
tx.execute(
"INSERT INTO memory_tags (memory_key, tag) VALUES (?1, ?2)",
params![key, tag],
)?;
}
tx.commit()?;
self.get_memory(key)?
.ok_or(rusqlite::Error::QueryReturnedNoRows)
}
pub fn delete_memory(&mut self, key: &str) -> Result<bool> {
let tx = self.conn.transaction()?;
tx.execute(
"DELETE FROM memory_tags WHERE memory_key = ?1",
params![key],
)?;
let affected = tx.execute("DELETE FROM memories WHERE key = ?1", params![key])?;
tx.commit()?;
Ok(affected > 0)
}
pub fn find_memories(
&self,
keyword: Option<&str>,
repo: Option<&str>,
tags: &[String],
limit: usize,
) -> Result<Vec<Memory>> {
let keyword_terms = normalize_keyword_terms(keyword);
let repo = normalize_repo(repo);
let tags = normalize_tags(tags);
let limit = limit.clamp(1, 50) as i64;
let keys = self.find_memory_keys(&keyword_terms, repo.as_deref(), &tags, limit)?;
keys.into_iter()
.map(|key| {
self.get_memory(&key)?
.ok_or(rusqlite::Error::QueryReturnedNoRows)
})
.collect()
}
fn get_memory(&self, key: &str) -> Result<Option<Memory>> {
self.conn
.query_row(
"
SELECT key, repo, content, created_at, updated_at
FROM memories
WHERE key = ?1
",
params![key],
|row| self.map_memory(row),
)
.optional()
}
fn find_memory_keys(
&self,
keyword_terms: &[String],
repo: Option<&str>,
tags: &[String],
limit: i64,
) -> Result<Vec<String>> {
let mut sql = String::from("SELECT m.key FROM memories m WHERE 1 = 1");
let mut params = Vec::<&dyn ToSql>::new();
let patterns = keyword_terms
.iter()
.map(|term| format!("%{}%", escape_like(term)))
.collect::<Vec<_>>();
let repo_value = repo.map(str::to_owned);
let tag_count = tags.len() as i64;
if !patterns.is_empty() {
sql.push_str(" AND (");
for index in 0..patterns.len() {
if index > 0 {
sql.push_str(" OR ");
}
sql.push_str(
"
m.key LIKE ? ESCAPE '\\'
OR m.content LIKE ? ESCAPE '\\'
OR m.repo LIKE ? ESCAPE '\\'
OR EXISTS (
SELECT 1
FROM memory_tags mt
WHERE mt.memory_key = m.key
AND mt.tag LIKE ? ESCAPE '\\'
)
",
);
}
sql.push(')');
for pattern in &patterns {
params.push(pattern);
params.push(pattern);
params.push(pattern);
params.push(pattern);
}
}
if let Some(repo) = repo_value.as_ref() {
sql.push_str(" AND m.repo = ?");
params.push(repo);
}
if !tags.is_empty() {
sql.push_str(
"
AND m.key IN (
SELECT memory_key
FROM memory_tags
WHERE tag IN (",
);
append_placeholders(&mut sql, tags.len());
sql.push_str(
")
GROUP BY memory_key
HAVING COUNT(DISTINCT tag) = ?
)",
);
params.extend(tags.iter().map(|tag| tag as &dyn ToSql));
params.push(&tag_count);
}
sql.push_str(" ORDER BY m.updated_at DESC, m.key ASC LIMIT ?");
params.push(&limit);
let mut stmt = self.conn.prepare(&sql)?;
let rows = stmt.query_map(rusqlite::params_from_iter(params), |row| row.get(0))?;
rows.collect()
}
fn map_memory(&self, row: &rusqlite::Row<'_>) -> Result<Memory> {
let key = row.get::<_, String>(0)?;
Ok(Memory {
repo: row.get(1)?,
tags: self.get_tags(&key)?,
key,
content: row.get(2)?,
created_at: row.get(3)?,
updated_at: row.get(4)?,
})
}
fn get_tags(&self, key: &str) -> Result<Vec<String>> {
let mut stmt = self.conn.prepare(
"
SELECT tag
FROM memory_tags
WHERE memory_key = ?1
ORDER BY tag ASC
",
)?;
let rows = stmt.query_map(params![key], |row| row.get(0))?;
rows.collect()
}
}
fn append_placeholders(sql: &mut String, count: usize) {
for index in 0..count {
if index > 0 {
sql.push_str(", ");
}
sql.push('?');
}
}
fn normalize_tags(tags: &[String]) -> Vec<String> {
tags.iter()
.map(|tag| tag.trim().to_lowercase())
.filter(|tag| !tag.is_empty())
.collect::<BTreeSet<_>>()
.into_iter()
.collect()
}
fn normalize_repo(repo: Option<&str>) -> Option<String> {
repo.map(str::trim)
.filter(|repo| !repo.is_empty())
.map(str::to_lowercase)
}
fn normalize_keyword_terms(keyword: Option<&str>) -> Vec<String> {
keyword
.unwrap_or_default()
.split_whitespace()
.map(normalize_keyword_term)
.filter(|term| !term.is_empty())
.collect::<BTreeSet<_>>()
.into_iter()
.collect()
}
fn normalize_keyword_term(term: &str) -> String {
term.trim()
.trim_matches(|ch: char| ch.is_ascii_punctuation() && ch != '-' && ch != '_' && ch != '%')
.to_lowercase()
}
fn escape_like(keyword: &str) -> String {
keyword
.replace('\\', "\\\\")
.replace('%', "\\%")
.replace('_', "\\_")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn upsert_inserts_and_updates_memory() {
let mut db = Database::open_in_memory().unwrap();
let inserted = db
.upsert_memory("project-style", None, &[], "follow existing pattern")
.unwrap();
let updated = db
.upsert_memory("project-style", None, &[], "keep implementation simple")
.unwrap();
assert_eq!(inserted.key, "project-style");
assert_eq!(updated.key, "project-style");
assert_eq!(updated.content, "keep implementation simple");
assert_eq!(inserted.created_at, updated.created_at);
}
#[test]
fn find_memories_matches_key_or_content() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory("rust", None, &[], "sqlite memory store")
.unwrap();
db.upsert_memory("mcp", None, &[], "tool server").unwrap();
let by_key = db.find_memories(Some("rust"), None, &[], 10).unwrap();
let by_content = db.find_memories(Some("tool"), None, &[], 10).unwrap();
assert_eq!(by_key.len(), 1);
assert_eq!(by_key[0].key, "rust");
assert_eq!(by_content.len(), 1);
assert_eq!(by_content[0].key, "mcp");
}
#[test]
fn delete_memory_reports_whether_row_existed() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory("temp", None, &[], "delete me").unwrap();
assert!(db.delete_memory("temp").unwrap());
assert!(!db.delete_memory("temp").unwrap());
}
#[test]
fn find_memories_matches_tags() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory(
"rust-style",
None,
&["rust".into(), "style".into()],
"use fmt",
)
.unwrap();
db.upsert_memory(
"mcp-style",
None,
&["mcp".into(), "style".into()],
"tool schema",
)
.unwrap();
let by_one_tag = db.find_memories(None, None, &["rust".into()], 10).unwrap();
let by_all_tags = db
.find_memories(None, None, &["style".into(), "mcp".into()], 10)
.unwrap();
assert_eq!(by_one_tag.len(), 1);
assert_eq!(by_one_tag[0].key, "rust-style");
assert_eq!(by_one_tag[0].tags, vec!["rust", "style"]);
assert_eq!(by_all_tags.len(), 1);
assert_eq!(by_all_tags[0].key, "mcp-style");
}
#[test]
fn upsert_replaces_tags() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory("memory", None, &["old".into()], "value")
.unwrap();
db.upsert_memory("memory", None, &["new".into()], "value")
.unwrap();
assert!(
db.find_memories(None, None, &["old".into()], 10)
.unwrap()
.is_empty()
);
assert_eq!(
db.find_memories(None, None, &["new".into()], 10).unwrap()[0].key,
"memory"
);
}
#[test]
fn find_memories_matches_repo() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory(
"gateway-login",
Some("infinity-gateway"),
&["login".into()],
"sync contact from access",
)
.unwrap();
db.upsert_memory(
"kira-session",
Some("kira"),
&["session".into()],
"persist reasoning",
)
.unwrap();
let gateway = db
.find_memories(None, Some("infinity-gateway"), &[], 10)
.unwrap();
let kira = db
.find_memories(Some("persist"), Some("kira"), &[], 10)
.unwrap();
assert_eq!(gateway.len(), 1);
assert_eq!(gateway[0].key, "gateway-login");
assert_eq!(gateway[0].repo.as_deref(), Some("infinity-gateway"));
assert_eq!(kira.len(), 1);
assert_eq!(kira[0].key, "kira-session");
}
#[test]
fn find_memories_matches_repo_and_tags() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory(
"gateway-login",
Some("infinity-gateway"),
&["login".into(), "contact".into()],
"sync contact from access",
)
.unwrap();
db.upsert_memory(
"gateway-token",
Some("infinity-gateway"),
&["auth".into()],
"token refresh",
)
.unwrap();
db.upsert_memory(
"kira-contact",
Some("kira"),
&["contact".into()],
"unrelated repo",
)
.unwrap();
let memories = db
.find_memories(None, Some("infinity-gateway"), &["contact".into()], 10)
.unwrap();
assert_eq!(memories.len(), 1);
assert_eq!(memories[0].key, "gateway-login");
}
#[test]
fn find_memories_without_filters_returns_recent_context() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory(
"first-memory",
Some("repo-a"),
&["decision".into()],
"first context",
)
.unwrap();
db.upsert_memory(
"second-memory",
Some("repo-b"),
&["bug".into()],
"second context",
)
.unwrap();
let memories = db.find_memories(None, None, &[], 10).unwrap();
assert_eq!(memories.len(), 2);
}
#[test]
fn find_memories_matches_keyword_terms_across_content_repo_and_tags() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory(
"gateway-token",
Some("infinity-gateway"),
&["auth".into()],
"token refresh rule",
)
.unwrap();
db.upsert_memory(
"kira-session",
Some("kira"),
&["persistence".into()],
"reasoning must be restored",
)
.unwrap();
let memories = db
.find_memories(Some("gateway reasoning"), None, &[], 10)
.unwrap();
let keys = memories
.into_iter()
.map(|memory| memory.key)
.collect::<BTreeSet<_>>();
assert!(keys.contains("gateway-token"));
assert!(keys.contains("kira-session"));
}
#[test]
fn find_memories_matches_quoted_keyword_terms() {
let mut db = Database::open_in_memory().unwrap();
db.upsert_memory(
"duplicate-tags",
Some("mcp-pocket-memory"),
&["edge-case".into()],
"Memory dengan duplicate tags di array",
)
.unwrap();
let memories = db
.find_memories(Some("\"duplicate tags\""), None, &[], 10)
.unwrap();
assert_eq!(memories.len(), 1);
assert_eq!(memories[0].key, "duplicate-tags");
}
#[test]
fn file_connections_support_parallel_writes() {
let db_path = std::env::temp_dir().join(format!(
"mcp-pocket-memory-parallel-{}.sqlite3",
std::process::id()
));
let _ = std::fs::remove_file(&db_path);
Database::open(&db_path).unwrap();
std::thread::scope(|scope| {
for index in 0..8 {
let db_path = db_path.clone();
scope.spawn(move || {
let mut db = Database::open(&db_path).unwrap();
db.upsert_memory(
&format!("parallel-memory-{index}"),
Some("parallel-repo"),
&["parallel".into()],
&format!("parallel content {index}"),
)
.unwrap();
});
}
});
let db = Database::open(&db_path).unwrap();
let memories = db
.find_memories(None, Some("parallel-repo"), &[], 20)
.unwrap();
let _ = std::fs::remove_file(&db_path);
let _ = std::fs::remove_file(db_path.with_extension("sqlite3-wal"));
let _ = std::fs::remove_file(db_path.with_extension("sqlite3-shm"));
assert_eq!(memories.len(), 8);
}
}