use std::collections::HashMap;
use std::path::Path;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use cel_memory::{
CallerScope, MemoryChunk, MemoryError, MemoryProvider, MemoryQuery, MemorySession, MemoryStats,
MemoryTier, NewMemoryChunk, NewMemorySession, Result as MemoryResult, SessionFilter,
SessionOutcome,
};
use chrono::Utc;
use duckdb::{params, Connection};
use uuid::Uuid;
use crate::error::DuckdbMemoryError;
use crate::util::{
chunk_matches_query, duck_err, embedding_literal, ids_in_clause, kind_str, optional_row,
outcome_str, row_to_chunk, row_to_session, rrf, source_str, tier_str, EMBEDDING_DIM,
};
use cel_memory::Embedder;
const GET_CHUNK_SQL: &str =
"SELECT id, created_at, kind, tier, source, session_id, project_root, caller_id, content, \
CAST(metadata AS VARCHAR) AS metadata, importance, pinned, shareable, superseded_by, \
embedding_model, embedding_dim FROM memory_chunks WHERE id = ?";
const SESSION_SELECT: &str =
"SELECT id, started_at, ended_at, caller_id, title, summary, outcome, \
CAST(metadata AS VARCHAR) AS metadata FROM memory_sessions";
const CHUNK_SELECT: &str =
"SELECT id, created_at, kind, tier, source, session_id, project_root, caller_id, content, \
CAST(metadata AS VARCHAR) AS metadata, importance, pinned, shareable, superseded_by, \
embedding_model, embedding_dim FROM memory_chunks";
pub struct DuckdbMemoryProvider {
conn: Arc<Mutex<Connection>>,
embedder: Arc<dyn Embedder>,
write_hook: Option<Arc<dyn cel_memory::MemoryWriteHook>>,
summarizer: Option<Arc<dyn cel_memory::Summarizer>>,
}
impl DuckdbMemoryProvider {
pub async fn open(
path: impl AsRef<Path>,
embedder: Arc<dyn Embedder>,
) -> Result<Self, DuckdbMemoryError> {
let path = path.as_ref().to_path_buf();
if embedder.dim() != EMBEDDING_DIM {
return Err(DuckdbMemoryError::DimMismatch {
expected: EMBEDDING_DIM,
actual: embedder.dim(),
});
}
let conn = tokio::task::spawn_blocking(move || -> Result<Connection, DuckdbMemoryError> {
let conn = Connection::open(&path)?;
crate::migrations::run(&conn)?;
Ok(conn)
})
.await
.map_err(|e| DuckdbMemoryError::BlockingJoin(e.to_string()))??;
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
embedder,
write_hook: None,
summarizer: None,
})
}
pub async fn open_in_memory(embedder: Arc<dyn Embedder>) -> Result<Self, DuckdbMemoryError> {
if embedder.dim() != EMBEDDING_DIM {
return Err(DuckdbMemoryError::DimMismatch {
expected: EMBEDDING_DIM,
actual: embedder.dim(),
});
}
let conn = tokio::task::spawn_blocking(|| -> Result<Connection, DuckdbMemoryError> {
let conn = Connection::open_in_memory()?;
crate::migrations::run(&conn)?;
Ok(conn)
})
.await
.map_err(|e| DuckdbMemoryError::BlockingJoin(e.to_string()))??;
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
embedder,
write_hook: None,
summarizer: None,
})
}
pub fn with_write_hook(mut self, hook: Arc<dyn cel_memory::MemoryWriteHook>) -> Self {
self.write_hook = Some(hook);
self
}
pub fn with_summarizer(mut self, summarizer: Arc<dyn cel_memory::Summarizer>) -> Self {
self.summarizer = Some(summarizer);
self
}
#[doc(hidden)]
pub fn conn_for_test(&self) -> Arc<Mutex<Connection>> {
Arc::clone(&self.conn)
}
async fn run<F, T>(&self, f: F) -> MemoryResult<T>
where
F: FnOnce(&Connection) -> MemoryResult<T> + Send + 'static,
T: Send + 'static,
{
let conn = Arc::clone(&self.conn);
tokio::task::spawn_blocking(move || {
let guard = conn
.lock()
.map_err(|e| MemoryError::Storage(format!("duckdb lock poisoned: {e}")))?;
f(&guard)
})
.await
.map_err(|e| MemoryError::Storage(format!("blocking join failed: {e}")))?
}
}
#[async_trait]
impl MemoryProvider for DuckdbMemoryProvider {
async fn retrieve(&self, query: MemoryQuery) -> MemoryResult<Vec<MemoryChunk>> {
if query.text.trim().is_empty() {
return Err(MemoryError::InvalidArgument(
"query.text must not be empty".into(),
));
}
let k = query.k.max(1);
let candidate_k = (3 * k).max(16);
let q_embedding = self.embedder.embed(&query.text).await?;
if q_embedding.len() != EMBEDDING_DIM {
return Err(MemoryError::Internal(format!(
"embedder produced dim {}, expected {EMBEDDING_DIM}",
q_embedding.len()
)));
}
let q_lit = embedding_literal(&q_embedding);
let scope = query.caller_scope;
let caller = query.caller_id.clone();
let text = query.text.clone();
let query_filter = query.clone();
self.run(move |conn| {
let vector_sql = match scope {
CallerScope::Global => format!(
"SELECT c.id
FROM memory_chunks c
INNER JOIN memory_vectors v ON v.chunk_id = c.id
ORDER BY array_cosine_distance(v.embedding, {q_lit}::FLOAT[{EMBEDDING_DIM}])
LIMIT {candidate_k}"
),
CallerScope::Own => format!(
"SELECT c.id
FROM memory_chunks c
INNER JOIN memory_vectors v ON v.chunk_id = c.id
WHERE c.caller_id = ?
ORDER BY array_cosine_distance(v.embedding, {q_lit}::FLOAT[{EMBEDDING_DIM}])
LIMIT {candidate_k}"
),
CallerScope::OwnPlusShared => format!(
"SELECT c.id
FROM memory_chunks c
INNER JOIN memory_vectors v ON v.chunk_id = c.id
WHERE (c.caller_id = ? OR c.shareable = TRUE)
ORDER BY array_cosine_distance(v.embedding, {q_lit}::FLOAT[{EMBEDDING_DIM}])
LIMIT {candidate_k}"
),
};
let vector_ids: Vec<String> = match scope {
CallerScope::Global => conn
.prepare(&vector_sql)
.map_err(duck_err)?
.query_map([], |row| row.get(0))
.map_err(duck_err)?
.collect::<Result<_, _>>()
.map_err(duck_err)?,
CallerScope::Own | CallerScope::OwnPlusShared => conn
.prepare(&vector_sql)
.map_err(duck_err)?
.query_map(params![caller.as_str()], |row| row.get(0))
.map_err(duck_err)?
.collect::<Result<_, _>>()
.map_err(duck_err)?,
};
let lexical_sql = match scope {
CallerScope::Global => {
"SELECT c.id
FROM memory_chunks c
WHERE c.content ILIKE '%' || ? || '%'
ORDER BY length(c.content)
LIMIT ?"
}
CallerScope::Own => {
"SELECT c.id
FROM memory_chunks c
WHERE c.caller_id = ?
AND c.content ILIKE '%' || ? || '%'
ORDER BY length(c.content)
LIMIT ?"
}
CallerScope::OwnPlusShared => {
"SELECT c.id
FROM memory_chunks c
WHERE (c.caller_id = ? OR c.shareable = TRUE)
AND c.content ILIKE '%' || ? || '%'
ORDER BY length(c.content)
LIMIT ?"
}
};
let lexical_ids: Vec<String> = match scope {
CallerScope::Global => conn
.prepare(lexical_sql)
.map_err(duck_err)?
.query_map(params![text.as_str(), candidate_k as i64], |row| row.get(0))
.map_err(duck_err)?
.collect::<Result<_, _>>()
.map_err(duck_err)?,
CallerScope::Own | CallerScope::OwnPlusShared => conn
.prepare(lexical_sql)
.map_err(duck_err)?
.query_map(
params![caller.as_str(), text.as_str(), candidate_k as i64],
|row| row.get(0),
)
.map_err(duck_err)?
.collect::<Result<_, _>>()
.map_err(duck_err)?,
};
let mut scores: HashMap<String, f32> = HashMap::new();
for (rank, id) in vector_ids.into_iter().enumerate() {
*scores.entry(id).or_insert(0.0) += rrf(rank, 60.0);
}
for (rank, id) in lexical_ids.into_iter().enumerate() {
*scores.entry(id).or_insert(0.0) += rrf(rank, 60.0);
}
if scores.is_empty() {
return Ok(Vec::new());
}
let ids: Vec<String> = scores.keys().cloned().collect();
let select_sql = format!(
"{CHUNK_SELECT} WHERE id IN ({})",
ids_in_clause(&ids)
);
let mut stmt = conn.prepare(&select_sql).map_err(duck_err)?;
let rows = stmt
.query_map([], row_to_chunk)
.map_err(duck_err)?;
let mut filtered = Vec::new();
for row in rows {
let c = row.map_err(duck_err)?;
if chunk_matches_query(&c, &query_filter) {
filtered.push(c);
}
}
filtered.sort_by(|a, b| {
let sa = scores.get(&a.id).copied().unwrap_or(0.0);
let sb = scores.get(&b.id).copied().unwrap_or(0.0);
sb.partial_cmp(&sa).unwrap_or(std::cmp::Ordering::Equal)
});
filtered.truncate(k);
Ok(filtered)
})
.await
}
async fn get(&self, chunk_id: &str) -> MemoryResult<Option<MemoryChunk>> {
let chunk_id = chunk_id.to_string();
self.run(move |conn| {
optional_row(
conn.query_row(GET_CHUNK_SQL, params![chunk_id.as_str()], row_to_chunk),
Ok,
)
})
.await
}
async fn get_session(&self, session_id: &str) -> MemoryResult<Option<MemorySession>> {
let session_id = session_id.to_string();
self.run(move |conn| {
optional_row(
conn.query_row(
&format!("{SESSION_SELECT} WHERE id = ?"),
params![session_id.as_str()],
row_to_session,
),
Ok,
)
})
.await
}
async fn list_sessions(&self, filter: SessionFilter) -> MemoryResult<Vec<MemorySession>> {
self.run(move |conn| list_sessions_inner(conn, filter))
.await
}
async fn write(&self, new_chunk: NewMemoryChunk) -> MemoryResult<MemoryChunk> {
if new_chunk.content.trim().is_empty() {
return Err(MemoryError::InvalidArgument(
"content must not be empty".into(),
));
}
if let Some(hook) = &self.write_hook {
match hook.before_write(&new_chunk).await? {
cel_memory::WriteDecision::Allow => {}
cel_memory::WriteDecision::Redact { reason } => {
return Ok(MemoryChunk {
id: Uuid::now_v7().to_string(),
created_at: Utc::now(),
kind: new_chunk.kind,
tier: MemoryTier::Session,
source: new_chunk.source,
session_id: new_chunk.session_id,
project_root: new_chunk.project_root,
caller_id: new_chunk.caller_id,
content: format!("<redacted: {reason}>"),
metadata: serde_json::json!({"redacted": true, "reason": reason}),
importance: 0.0,
pinned: false,
shareable: false,
superseded_by: None,
embedding_model: "none".into(),
embedding_dim: 0,
});
}
}
}
let embedding = self.embedder.embed(&new_chunk.content).await?;
if embedding.len() != EMBEDDING_DIM {
return Err(MemoryError::Internal(format!(
"embedder produced dim {}, expected {EMBEDDING_DIM}",
embedding.len()
)));
}
let chunk = MemoryChunk {
id: Uuid::now_v7().to_string(),
created_at: Utc::now(),
kind: new_chunk.kind,
tier: MemoryTier::Session,
source: new_chunk.source,
session_id: new_chunk.session_id.clone(),
project_root: new_chunk.project_root.clone(),
caller_id: new_chunk.caller_id.clone(),
content: new_chunk.content.clone(),
metadata: new_chunk.metadata.clone(),
importance: cel_memory::score_importance(&new_chunk),
pinned: new_chunk.pinned,
shareable: new_chunk.shareable,
superseded_by: None,
embedding_model: self.embedder.model_name().to_string(),
embedding_dim: EMBEDDING_DIM as u32,
};
let metadata_json = if chunk.metadata.is_null() {
"{}".to_string()
} else {
chunk.metadata.to_string()
};
let emb_lit = embedding_literal(&embedding);
let chunk_id = chunk.id.clone();
self.run(move |conn| {
conn.execute_batch("BEGIN TRANSACTION")
.map_err(|e| MemoryError::Storage(e.to_string()))?;
let result = (|| -> MemoryResult<()> {
conn.execute(
"INSERT INTO memory_chunks(
id, created_at, kind, tier, source, session_id, project_root,
caller_id, content, metadata, importance, pinned, shareable,
superseded_by, embedding_model, embedding_dim
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
params![
chunk.id.as_str(),
chunk.created_at,
kind_str(chunk.kind),
tier_str(chunk.tier),
source_str(chunk.source),
chunk.session_id.as_deref(),
chunk.project_root.as_deref(),
chunk.caller_id.as_str(),
chunk.content.as_str(),
metadata_json.as_str(),
chunk.importance,
chunk.pinned,
chunk.shareable,
chunk.superseded_by.as_deref(),
chunk.embedding_model.as_str(),
chunk.embedding_dim as i32,
],
)
.map_err(|e| MemoryError::Storage(e.to_string()))?;
conn.execute_batch(&format!(
"INSERT INTO memory_vectors (chunk_id, embedding) VALUES ({}, {}::FLOAT[{EMBEDDING_DIM}])",
crate::util::sql_string_literal(&chunk_id),
emb_lit
))
.map_err(|e| MemoryError::Storage(e.to_string()))?;
Ok(())
})();
if result.is_ok() {
conn.execute_batch("COMMIT")
.map_err(|e| MemoryError::Storage(e.to_string()))?;
} else {
conn.execute_batch("ROLLBACK").ok();
result?;
}
Ok(chunk)
})
.await
}
async fn write_batch(&self, chunks: Vec<NewMemoryChunk>) -> MemoryResult<Vec<MemoryChunk>> {
let mut out = Vec::with_capacity(chunks.len());
for chunk in chunks {
out.push(self.write(chunk).await?);
}
Ok(out)
}
async fn open_session(&self, init: NewMemorySession) -> MemoryResult<MemorySession> {
let session = MemorySession {
id: Uuid::now_v7().to_string(),
started_at: Utc::now(),
ended_at: None,
caller_id: init.caller_id,
title: init.title,
summary: None,
outcome: SessionOutcome::Open,
metadata: init.metadata,
};
let metadata_json = if session.metadata.is_null() {
"{}".to_string()
} else {
session.metadata.to_string()
};
let session_id = session.id.clone();
self.run(move |conn| {
conn.execute(
"INSERT INTO memory_sessions (id, started_at, ended_at, caller_id, title, summary, outcome, metadata)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
params![
session_id.as_str(),
session.started_at,
session.ended_at,
session.caller_id.as_str(),
session.title.as_deref(),
session.summary.as_deref(),
outcome_str(session.outcome),
metadata_json.as_str(),
],
)
.map_err(|e| MemoryError::Storage(e.to_string()))?;
Ok(session)
})
.await
}
async fn close_session(&self, session_id: &str, outcome: SessionOutcome) -> MemoryResult<()> {
let session_id = session_id.to_string();
self.run(move |conn| {
let changed = conn
.execute(
"UPDATE memory_sessions SET ended_at = ?, outcome = ? WHERE id = ?",
params![Utc::now(), outcome_str(outcome), session_id.as_str()],
)
.map_err(|e| MemoryError::Storage(e.to_string()))?;
if changed == 0 {
return Err(MemoryError::NotFound(session_id));
}
Ok(())
})
.await
}
async fn rename_session(&self, session_id: &str, title: &str) -> MemoryResult<()> {
let session_id = session_id.to_string();
let title = title.to_string();
self.run(move |conn| {
let changed = conn
.execute(
"UPDATE memory_sessions SET title = ? WHERE id = ?",
params![title.as_str(), session_id.as_str()],
)
.map_err(|e| MemoryError::Storage(e.to_string()))?;
if changed == 0 {
return Err(MemoryError::NotFound(session_id));
}
Ok(())
})
.await
}
async fn stats(&self) -> MemoryResult<MemoryStats> {
self.run(|conn| {
let total_chunks: i64 = conn
.query_row("SELECT COUNT(*)::BIGINT FROM memory_chunks", [], |row| {
row.get(0)
})
.map_err(|e| MemoryError::Storage(e.to_string()))?;
let total_sessions: i64 = conn
.query_row("SELECT COUNT(*)::BIGINT FROM memory_sessions", [], |row| {
row.get(0)
})
.map_err(|e| MemoryError::Storage(e.to_string()))?;
let embedding_model = optional_row(
conn.query_row(
"SELECT embedding_model FROM memory_chunks ORDER BY created_at DESC LIMIT 1",
[],
|row| row.get::<_, String>(0),
),
Ok,
)?;
Ok(MemoryStats {
total_chunks: total_chunks as usize,
total_sessions: total_sessions as usize,
embedding_model,
..MemoryStats::default()
})
})
.await
}
async fn summarize_session(&self, _session_id: &str) -> MemoryResult<MemoryChunk> {
if self.summarizer.is_none() {
return Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::summarize_session — attach summarizer via with_summarizer",
));
}
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::summarize_session — Phase 3",
))
}
async fn rollup_day(&self, _date: chrono::NaiveDate) -> MemoryResult<Vec<MemoryChunk>> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::rollup_day — Phase 3",
))
}
async fn rollup_rule_week(
&self,
_rule_id: &str,
_week_start: chrono::NaiveDate,
) -> MemoryResult<MemoryChunk> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::rollup_rule_week — Phase 3",
))
}
async fn run_aging_sweep(&self) -> MemoryResult<cel_memory::AgingReport> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::run_aging_sweep — Phase 2",
))
}
async fn re_embed_all(&self, _target_model: &str) -> MemoryResult<cel_memory::ReEmbedReport> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::re_embed_all — Phase 4",
))
}
async fn export(
&self,
_filter: cel_memory::ExportFilter,
) -> MemoryResult<cel_memory::ExportBundle> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::export — Phase 2",
))
}
async fn pin(&self, _chunk_id: &str, _pinned: bool) -> MemoryResult<()> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::pin — Phase 2",
))
}
async fn update_importance(&self, _chunk_id: &str, _importance: f32) -> MemoryResult<()> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::update_importance — Phase 2",
))
}
async fn supersede(&self, _old_id: &str, _new_id: &str) -> MemoryResult<()> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::supersede — Phase 2",
))
}
async fn record_access(
&self,
_chunk_id: &str,
_retrieved_by: &str,
_used: bool,
) -> MemoryResult<()> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::record_access — Phase 2",
))
}
async fn delete(
&self,
_chunk_id: &str,
_reason: cel_memory::EvictionReason,
) -> MemoryResult<()> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::delete — Phase 2",
))
}
async fn delete_matching(
&self,
_predicate: cel_memory::MemoryPredicate,
_reason: cel_memory::EvictionReason,
) -> MemoryResult<usize> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::delete_matching — Phase 2",
))
}
async fn purge_all(&self) -> MemoryResult<cel_memory::PurgeReport> {
Err(MemoryError::NotImplemented(
"DuckdbMemoryProvider::purge_all — Phase 2",
))
}
}
fn list_sessions_inner(
conn: &Connection,
filter: SessionFilter,
) -> MemoryResult<Vec<MemorySession>> {
let mut out = Vec::new();
match (&filter.caller_id, filter.open_only, filter.outcome) {
(None, false, None) => {
let mut stmt = conn
.prepare(&format!("{SESSION_SELECT} ORDER BY started_at DESC"))
.map_err(duck_err)?;
for row in stmt
.query_map([], row_to_session)
.map_err(duck_err)? {
out.push(row.map_err(duck_err)?);
}
}
(Some(caller), false, None) => {
let mut stmt = conn
.prepare(&format!(
"{SESSION_SELECT} WHERE caller_id = ? ORDER BY started_at DESC"
))
.map_err(duck_err)?;
for row in stmt
.query_map(params![caller.as_str()], row_to_session)
.map_err(duck_err)?
{
out.push(row.map_err(duck_err)?);
}
}
(None, true, _) => {
let mut stmt = conn
.prepare(&format!(
"{SESSION_SELECT} WHERE outcome = 'open' ORDER BY started_at DESC"
))
.map_err(duck_err)?;
for row in stmt
.query_map([], row_to_session)
.map_err(duck_err)? {
out.push(row.map_err(duck_err)?);
}
}
(Some(caller), true, _) => {
let mut stmt = conn
.prepare(&format!(
"{SESSION_SELECT} WHERE caller_id = ? AND outcome = 'open' ORDER BY started_at DESC"
))
.map_err(duck_err)?;
for row in stmt
.query_map(params![caller.as_str()], row_to_session)
.map_err(duck_err)?
{
out.push(row.map_err(duck_err)?);
}
}
(None, false, Some(outcome)) => {
let mut stmt = conn
.prepare(&format!(
"{SESSION_SELECT} WHERE outcome = ? ORDER BY started_at DESC"
))
.map_err(duck_err)?;
for row in stmt
.query_map(params![outcome_str(outcome)], row_to_session)
.map_err(duck_err)?
{
out.push(row.map_err(duck_err)?);
}
}
(Some(caller), false, Some(outcome)) => {
let mut stmt = conn
.prepare(&format!(
"{SESSION_SELECT} WHERE caller_id = ? AND outcome = ? ORDER BY started_at DESC"
))
.map_err(duck_err)?;
for row in stmt
.query_map(params![caller.as_str(), outcome_str(outcome)], row_to_session)
.map_err(duck_err)?
{
out.push(row.map_err(duck_err)?);
}
}
}
Ok(out)
}