use async_trait::async_trait;
use hnsw_rs::prelude::{DistCosine, Hnsw};
use meerkat_core::memory::{
EmbeddingModel, HnswParams, MemoryEnumerationPage, MemoryEnumerationRequest, MemoryIndexBatch,
MemoryIndexReceipt, MemoryMetadata, MemoryOwner, MemoryRankingPolicy, MemoryRecord,
MemoryResult, MemoryScopeDropReceipt, MemorySearchScope, MemoryStore, MemoryStoreError,
};
use meerkat_core::types::SessionId;
use rusqlite::{Connection, OptionalExtension, Transaction, TransactionBehavior, params};
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
const SQLITE_BUSY_TIMEOUT_MS: u64 = 5_000;
const CREATE_MEMORY_SCHEMA_SQL: &str = r"
CREATE TABLE IF NOT EXISTS memory_metadata (
point_id INTEGER PRIMARY KEY,
metadata_json BLOB NOT NULL
);
CREATE TABLE IF NOT EXISTS memory_text (
point_id INTEGER PRIMARY KEY,
content BLOB NOT NULL
)";
const CREATE_MEMORY_SESSION_INDEX_SQL: &str =
"CREATE INDEX IF NOT EXISTS idx_memory_metadata_session_id ON memory_metadata (session_id)";
const CREATE_MEMORY_ALLOCATOR_SCHEMA_SQL: &str = "CREATE TABLE IF NOT EXISTS memory_allocator (
id INTEGER PRIMARY KEY CHECK (id = 0),
high_water INTEGER NOT NULL
)";
const SEED_MEMORY_ALLOCATOR_SQL: &str = "INSERT INTO memory_allocator (id, high_water) \
SELECT 0, COALESCE((SELECT MAX(point_id) + 1 FROM memory_metadata), 0) \
WHERE NOT EXISTS (SELECT 1 FROM memory_allocator WHERE id = 0)";
const DEFAULT_VOCAB_DIM: usize = 4096;
const MIN_INDEX_ELEMENTS_HINT: usize = 1;
pub fn default_ranking_policy() -> MemoryRankingPolicy {
MemoryRankingPolicy::new(
Arc::new(BagOfWordsEmbeddingModel::new(DEFAULT_VOCAB_DIM)),
HnswParams::default(),
)
}
pub struct BagOfWordsEmbeddingModel {
vocab_dim: usize,
}
impl BagOfWordsEmbeddingModel {
pub fn new(vocab_dim: usize) -> Self {
Self { vocab_dim }
}
}
impl EmbeddingModel for BagOfWordsEmbeddingModel {
fn dimension(&self) -> usize {
self.vocab_dim
}
fn embed(&self, text: &str) -> Vec<f32> {
let mut vec = vec![0.0f32; self.vocab_dim];
for word in text.split_whitespace() {
let hash = word.bytes().fold(0usize, |acc, b| {
acc.wrapping_mul(31)
.wrapping_add(b.to_ascii_lowercase() as usize)
}) % self.vocab_dim;
vec[hash] += 1.0;
}
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut vec {
*x /= norm;
}
}
vec
}
}
fn open_connection(path: &Path) -> Result<Connection, MemoryStoreError> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).map_err(MemoryStoreError::Io)?;
}
let conn = Connection::open(path).map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
conn.busy_timeout(Duration::from_millis(SQLITE_BUSY_TIMEOUT_MS))
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
conn.pragma_update(None, "journal_mode", "WAL")
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
conn.pragma_update(None, "synchronous", "FULL")
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
conn.execute_batch(CREATE_MEMORY_SCHEMA_SQL)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
Ok(conn)
}
fn migrate_memory_schema(conn: &mut Connection) -> Result<(), MemoryStoreError> {
let tx = conn
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let has_session_id = {
let mut stmt = tx
.prepare("PRAGMA table_info(memory_metadata)")
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let column_names = stmt
.query_map([], |row| row.get::<_, String>(1))
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let mut found = false;
for name in column_names {
let name = name.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
if name == "session_id" {
found = true;
}
}
found
};
if !has_session_id {
tx.execute("ALTER TABLE memory_metadata ADD COLUMN session_id TEXT", [])
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
backfill_null_session_ids(&tx)?;
}
tx.execute(CREATE_MEMORY_SESSION_INDEX_SQL, [])
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
tx.execute(CREATE_MEMORY_ALLOCATOR_SCHEMA_SQL, [])
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
tx.execute(SEED_MEMORY_ALLOCATOR_SQL, [])
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
tx.commit()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
Ok(())
}
fn backfill_null_session_ids(conn: &Connection) -> Result<(), MemoryStoreError> {
let mut null_rows = Vec::new();
{
let mut stmt = conn
.prepare("SELECT point_id, metadata_json FROM memory_metadata WHERE session_id IS NULL")
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let mapped = stmt
.query_map([], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, Vec<u8>>(1)?))
})
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
for row in mapped {
null_rows.push(row.map_err(|e| MemoryStoreError::Storage(e.to_string()))?);
}
}
for (point_id, metadata_json) in null_rows {
let metadata: MemoryMetadata = serde_json::from_slice(&metadata_json)
.map_err(|e| MemoryStoreError::Embedding(e.to_string()))?;
conn.execute(
"UPDATE memory_metadata SET session_id = ?1 WHERE point_id = ?2",
params![metadata.session_id.to_string(), point_id],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
}
Ok(())
}
fn allocate_point_ids(tx: &Transaction<'_>, count: usize) -> Result<Vec<i64>, MemoryStoreError> {
let high_water: i64 = tx
.query_row(
"SELECT high_water FROM memory_allocator WHERE id = 0",
[],
|row| row.get(0),
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let max_row: Option<i64> = tx
.query_row("SELECT MAX(point_id) FROM memory_metadata", [], |row| {
row.get(0)
})
.optional()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?
.flatten();
let row_floor = match max_row {
Some(max) => max
.checked_add(1)
.ok_or(MemoryStoreError::PointIdOverflow)?,
None => 0,
};
let base = high_water.max(row_floor);
let count = i64::try_from(count).map_err(|_| MemoryStoreError::PointIdOutOfRange)?;
let next_high = base
.checked_add(count)
.ok_or(MemoryStoreError::PointIdOverflow)?;
tx.execute(
"UPDATE memory_allocator SET high_water = ?1 WHERE id = 0",
params![next_high],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
Ok((base..next_high).collect())
}
type MemoryHnswIndex = Hnsw<'static, f32, DistCosine>;
fn bounded_index_elements_hint(entries: usize) -> usize {
entries.max(MIN_INDEX_ELEMENTS_HINT)
}
fn new_hnsw_index(max_elements_hint: usize, params: HnswParams) -> MemoryHnswIndex {
Hnsw::<'static, f32, DistCosine>::new(
params.max_nb_connection,
bounded_index_elements_hint(max_elements_hint),
params.max_layer,
params.ef_construction,
DistCosine {},
)
}
fn rebuild_scoped_index_from_db(
conn: &Connection,
session_id: &SessionId,
embedding_model: &dyn EmbeddingModel,
params: HnswParams,
) -> Result<ScopedHnswIndex, MemoryStoreError> {
backfill_null_session_ids(conn)?;
let mut scoped: Vec<(i64, Vec<u8>)> = Vec::new();
{
let mut stmt = conn
.prepare(
"SELECT t.point_id, t.content \
FROM memory_text t \
JOIN memory_metadata m ON m.point_id = t.point_id \
WHERE m.session_id = ?1 \
ORDER BY t.point_id ASC",
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let mapped = stmt
.query_map(params![session_id.to_string()], |row| {
Ok((row.get::<_, i64>(0)?, row.get::<_, Vec<u8>>(1)?))
})
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
for row in mapped {
scoped.push(row.map_err(|e| MemoryStoreError::Storage(e.to_string()))?);
}
}
let index = ScopedHnswIndex::new(scoped.len(), params);
for (point_id, text) in scoped {
let text = decode_memory_text(point_id, text)?;
let point_id =
usize::try_from(point_id).map_err(|_| MemoryStoreError::PointIdOutOfRange)?;
let embedding = embedding_model.embed(&text);
index.insert(&embedding, point_id);
}
Ok(index)
}
fn decode_memory_text(point_id: i64, bytes: Vec<u8>) -> Result<String, MemoryStoreError> {
String::from_utf8(bytes).map_err(|_| MemoryStoreError::TextCorruption { point_id })
}
enum ScopedIndexState {
Live(ScopedHnswIndex),
Poisoned,
}
struct ScopedHnswIndex {
index: MemoryHnswIndex,
#[cfg(test)]
max_elements_hint: usize,
}
impl ScopedHnswIndex {
fn new(max_elements_hint: usize, params: HnswParams) -> Self {
let max_elements_hint = bounded_index_elements_hint(max_elements_hint);
Self {
index: new_hnsw_index(max_elements_hint, params),
#[cfg(test)]
max_elements_hint,
}
}
fn insert(&self, embedding: &[f32], point_id: usize) {
self.index.insert((embedding, point_id));
}
}
pub struct HnswMemoryStore {
indices: Arc<std::sync::RwLock<HashMap<SessionId, ScopedIndexState>>>,
db_path: PathBuf,
insert_lock: Mutex<()>,
path: PathBuf,
policy: MemoryRankingPolicy,
#[cfg(test)]
fail_hnsw_insert_after: Arc<std::sync::atomic::AtomicI64>,
}
impl std::fmt::Debug for HnswMemoryStore {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HnswMemoryStore")
.field("db_path", &self.db_path)
.field("path", &self.path)
.finish_non_exhaustive()
}
}
impl HnswMemoryStore {
pub fn open(dir: impl AsRef<Path>) -> Result<Self, MemoryStoreError> {
Self::open_with_policy(dir, default_ranking_policy())
}
pub fn open_with_policy(
dir: impl AsRef<Path>,
policy: MemoryRankingPolicy,
) -> Result<Self, MemoryStoreError> {
let dir = dir.as_ref();
std::fs::create_dir_all(dir).map_err(MemoryStoreError::Io)?;
let db_path = dir.join("memory.sqlite3");
let mut conn = open_connection(&db_path)?;
migrate_memory_schema(&mut conn)?;
backfill_null_session_ids(&conn)?;
Ok(Self {
indices: Arc::new(std::sync::RwLock::new(HashMap::new())),
db_path,
insert_lock: Mutex::new(()),
path: dir.to_path_buf(),
policy,
#[cfg(test)]
fail_hnsw_insert_after: Arc::new(std::sync::atomic::AtomicI64::new(-1)),
})
}
pub fn path(&self) -> &Path {
&self.path
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
fn hnsw_index_count(&self) -> usize {
self.indices.read().unwrap().len()
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
fn hnsw_point_count(&self) -> usize {
self.indices
.read()
.unwrap()
.values()
.map(|state| match state {
ScopedIndexState::Live(index) => index.index.get_nb_point(),
ScopedIndexState::Poisoned => 0,
})
.sum()
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
fn hnsw_index_hints(&self) -> Vec<usize> {
self.indices
.read()
.unwrap()
.values()
.filter_map(|state| match state {
ScopedIndexState::Live(index) => Some(index.max_elements_hint),
ScopedIndexState::Poisoned => None,
})
.collect()
}
#[cfg(test)]
fn arm_hnsw_insert_failure_after(&self, after: i64) {
self.fail_hnsw_insert_after
.store(after, std::sync::atomic::Ordering::Release);
}
}
#[async_trait]
impl MemoryStore for HnswMemoryStore {
async fn index_scoped_batch(
&self,
batch: MemoryIndexBatch,
) -> Result<MemoryIndexReceipt, MemoryStoreError> {
let (receipt_scope, requests) = batch.into_parts();
let mut entries = Vec::with_capacity(requests.len());
for request in requests {
let (_scope, content, metadata) = request.into_parts();
if !content.is_indexable() {
continue;
}
let text = content.into_indexable_text();
let meta_json = serde_json::to_vec(&metadata)
.map_err(|e| MemoryStoreError::Embedding(e.to_string()))?;
let embedding = self.policy.embed(&text);
entries.push((text, meta_json, embedding));
}
let indexed_entries = entries.len();
if indexed_entries == 0 {
return Ok(MemoryIndexReceipt {
scope: receipt_scope,
indexed_entries: 0,
});
}
let db_path = self.db_path.clone();
let indices = Arc::clone(&self.indices);
let session_id = receipt_scope.session_id().clone();
let hnsw_params = self.policy.hnsw_params();
let embedding_model = Arc::clone(self.policy.embedding_model());
#[cfg(test)]
let fail_hnsw_insert_after = Arc::clone(&self.fail_hnsw_insert_after);
let _guard = self.insert_lock.lock().await;
let insert_result = tokio::task::spawn_blocking(move || {
let mut conn = open_connection(&db_path)?;
let tx = conn
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let point_ids = allocate_point_ids(&tx, indexed_entries)?;
let session_param = session_id.to_string();
for (point_id_i64, (content, meta_json, _embedding)) in point_ids.iter().zip(&entries) {
tx.execute(
"INSERT INTO memory_metadata (point_id, metadata_json, session_id) \
VALUES (?1, ?2, ?3)",
params![point_id_i64, meta_json, session_param],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
tx.execute(
"INSERT INTO memory_text (point_id, content) VALUES (?1, ?2)",
params![point_id_i64, content.as_bytes()],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
}
tx.commit()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let index_result = (|| {
let mut indices = indices
.write()
.map_err(|_| MemoryStoreError::LockPoisoned)?;
if matches!(indices.get(&session_id), Some(ScopedIndexState::Poisoned)) {
let rebuilt = rebuild_scoped_index_from_db(
&conn,
&session_id,
embedding_model.as_ref(),
hnsw_params,
)?;
indices.insert(session_id.clone(), ScopedIndexState::Live(rebuilt));
return Ok(());
}
let state = match indices.entry(session_id.clone()) {
std::collections::hash_map::Entry::Occupied(entry) => entry.into_mut(),
std::collections::hash_map::Entry::Vacant(entry) => {
let rebuilt = rebuild_scoped_index_from_db(
&conn,
&session_id,
embedding_model.as_ref(),
hnsw_params,
)?;
entry.insert(ScopedIndexState::Live(rebuilt));
return Ok(());
}
};
let ScopedIndexState::Live(index) = state else {
return Err(MemoryStoreError::ScopePoisoned);
};
#[allow(clippy::unused_enumerate_index)]
for (_ordinal, (point_id, (_content, _meta_json, embedding))) in
point_ids.iter().zip(&entries).enumerate()
{
#[cfg(test)]
{
let fail_after =
fail_hnsw_insert_after.load(std::sync::atomic::Ordering::Acquire);
if fail_after >= 0 && _ordinal as i64 >= fail_after {
return Err(MemoryStoreError::LockPoisoned);
}
}
let point_id = usize::try_from(*point_id)
.map_err(|_| MemoryStoreError::PointIdOutOfRange)?;
index.insert(embedding, point_id);
}
Ok::<(), MemoryStoreError>(())
})();
if let Err(error) = index_result {
let repair_result = (|| -> Result<ScopedHnswIndex, MemoryStoreError> {
let mut cleanup = open_connection(&db_path)?;
let tx = cleanup
.transaction()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
for point_id_i64 in &point_ids {
tx.execute(
"DELETE FROM memory_metadata WHERE point_id = ?1",
params![point_id_i64],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
tx.execute(
"DELETE FROM memory_text WHERE point_id = ?1",
params![point_id_i64],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
}
tx.commit()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
rebuild_scoped_index_from_db(
&cleanup,
&session_id,
embedding_model.as_ref(),
hnsw_params,
)
})();
let mut indices = indices
.write()
.map_err(|_| MemoryStoreError::LockPoisoned)?;
match repair_result {
Ok(repaired) => {
indices.insert(session_id, ScopedIndexState::Live(repaired));
return Err(error);
}
Err(repair) => {
indices.insert(session_id, ScopedIndexState::Poisoned);
return Err(MemoryStoreError::ScopeRepairFailed {
original: Box::new(error),
repair: Box::new(repair),
});
}
}
}
Ok::<(), MemoryStoreError>(())
})
.await
.map_err(|e| MemoryStoreError::TaskJoin(format!("index task join failed: {e}")))?;
insert_result?;
Ok(MemoryIndexReceipt {
scope: receipt_scope,
indexed_entries,
})
}
async fn search(
&self,
scope: &MemorySearchScope,
query: &str,
limit: usize,
) -> Result<Vec<MemoryResult>, MemoryStoreError> {
if limit == 0 {
return Ok(Vec::new());
}
match self.search_pass(scope, query, limit, false).await? {
Some(results) => Ok(results),
None => {
let _guard = self.insert_lock.lock().await;
match self.search_pass(scope, query, limit, true).await? {
Some(results) => Ok(results),
None => Err(MemoryStoreError::Storage(
"scope load pass did not publish a live index".to_string(),
)),
}
}
}
}
async fn enumerate_scoped(
&self,
scope: &MemorySearchScope,
request: MemoryEnumerationRequest,
) -> Result<MemoryEnumerationPage, MemoryStoreError> {
self.enumerate_scoped_impl(scope, request).await
}
async fn drop_scope(
&self,
owner: &MemoryOwner,
) -> Result<MemoryScopeDropReceipt, MemoryStoreError> {
self.drop_scope_impl(owner).await
}
}
impl HnswMemoryStore {
async fn search_pass(
&self,
scope: &MemorySearchScope,
query: &str,
limit: usize,
load_if_missing: bool,
) -> Result<Option<Vec<MemoryResult>>, MemoryStoreError> {
let query = query.to_owned();
let scope = scope.clone();
let db_path = self.db_path.clone();
let indices = Arc::clone(&self.indices);
let embedding_model = Arc::clone(self.policy.embedding_model());
let hnsw_params = self.policy.hnsw_params();
let ef_search = hnsw_params.ef_search;
tokio::task::spawn_blocking(move || {
let embedding = embedding_model.embed(&query);
let loaded_neighbors = {
let indices = indices.read().map_err(|_| MemoryStoreError::LockPoisoned)?;
match indices.get(scope.session_id()) {
None => None,
Some(ScopedIndexState::Poisoned) => {
return Err(MemoryStoreError::ScopePoisoned);
}
Some(ScopedIndexState::Live(index)) => {
Some(index.index.search(&embedding, limit, limit.max(ef_search)))
}
}
};
let conn = open_connection(&db_path)?;
let neighbors = match loaded_neighbors {
Some(neighbors) => neighbors,
None if !load_if_missing => return Ok(None),
None => {
let rebuilt = rebuild_scoped_index_from_db(
&conn,
scope.session_id(),
embedding_model.as_ref(),
hnsw_params,
)?;
let mut indices = indices
.write()
.map_err(|_| MemoryStoreError::LockPoisoned)?;
match indices
.entry(scope.session_id().clone())
.or_insert(ScopedIndexState::Live(rebuilt))
{
ScopedIndexState::Poisoned => {
return Err(MemoryStoreError::ScopePoisoned);
}
ScopedIndexState::Live(index) => {
index.index.search(&embedding, limit, limit.max(ef_search))
}
}
}
};
let mut results = Vec::with_capacity(neighbors.len());
for neighbor in &neighbors {
let point_id = i64::try_from(neighbor.d_id)
.map_err(|_| MemoryStoreError::PointIdOutOfRange)?;
let content = match conn
.query_row(
"SELECT content FROM memory_text WHERE point_id = ?1",
params![point_id],
|row| row.get::<_, Vec<u8>>(0),
)
.optional()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?
{
Some(bytes) => decode_memory_text(point_id, bytes)?,
None => return Err(MemoryStoreError::IndexDivergence { point_id }),
};
let metadata = match conn
.query_row(
"SELECT metadata_json FROM memory_metadata WHERE point_id = ?1",
params![point_id],
|row| row.get::<_, Vec<u8>>(0),
)
.optional()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?
{
Some(bytes) => serde_json::from_slice(&bytes)
.map_err(|e| MemoryStoreError::Embedding(e.to_string()))?,
None => return Err(MemoryStoreError::IndexDivergence { point_id }),
};
if !scope.includes(&metadata) {
continue;
}
let score = 1.0 - (neighbor.distance / 2.0);
results.push(MemoryResult {
content,
metadata,
score,
});
}
Ok::<Option<Vec<MemoryResult>>, MemoryStoreError>(Some(results))
})
.await
.map_err(|e| MemoryStoreError::TaskJoin(format!("search task join failed: {e}")))?
}
async fn drop_scope_impl(
&self,
owner: &MemoryOwner,
) -> Result<MemoryScopeDropReceipt, MemoryStoreError> {
let owner = owner.clone();
let db_path = self.db_path.clone();
let indices = Arc::clone(&self.indices);
let _guard = self.insert_lock.lock().await;
tokio::task::spawn_blocking(
move || -> Result<MemoryScopeDropReceipt, MemoryStoreError> {
let mut conn = open_connection(&db_path)?;
let tx = conn
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
backfill_null_session_ids(&tx)?;
let session_param = owner.session_id().to_string();
tx.execute(
"DELETE FROM memory_text WHERE point_id IN \
(SELECT point_id FROM memory_metadata WHERE session_id = ?1)",
params![session_param],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let dropped_entries = tx
.execute(
"DELETE FROM memory_metadata WHERE session_id = ?1",
params![session_param],
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
tx.commit()
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let mut indices = indices
.write()
.map_err(|_| MemoryStoreError::LockPoisoned)?;
indices.remove(owner.session_id());
Ok(MemoryScopeDropReceipt {
owner,
dropped_entries,
})
},
)
.await
.map_err(|e| MemoryStoreError::TaskJoin(format!("drop task join failed: {e}")))?
}
async fn enumerate_scoped_impl(
&self,
scope: &MemorySearchScope,
request: MemoryEnumerationRequest,
) -> Result<MemoryEnumerationPage, MemoryStoreError> {
if request.limit == 0 {
return Err(MemoryStoreError::EnumerationLimitZero);
}
let scope = scope.clone();
let db_path = self.db_path.clone();
tokio::task::spawn_blocking(move || -> Result<MemoryEnumerationPage, MemoryStoreError> {
let conn = open_connection(&db_path)?;
backfill_null_session_ids(&conn)?;
let fetch_limit = i64::try_from(request.limit.saturating_add(1)).unwrap_or(i64::MAX);
let fetch_offset = i64::try_from(request.offset).unwrap_or(i64::MAX);
let mut raw_rows: Vec<(i64, Option<Vec<u8>>, Vec<u8>)> = Vec::new();
{
let mut stmt = conn
.prepare(
"SELECT m.point_id, t.content, m.metadata_json \
FROM memory_metadata m \
LEFT JOIN memory_text t ON t.point_id = m.point_id \
WHERE m.session_id = ?1 \
ORDER BY m.point_id ASC \
LIMIT ?2 OFFSET ?3",
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
let mapped = stmt
.query_map(
params![scope.session_id().to_string(), fetch_limit, fetch_offset],
|row| {
Ok((
row.get::<_, i64>(0)?,
row.get::<_, Option<Vec<u8>>>(1)?,
row.get::<_, Vec<u8>>(2)?,
))
},
)
.map_err(|e| MemoryStoreError::Storage(e.to_string()))?;
for row in mapped {
raw_rows.push(row.map_err(|e| MemoryStoreError::Storage(e.to_string()))?);
}
}
let more_raw_rows_remain = raw_rows.len() > request.limit;
raw_rows.truncate(request.limit);
let rows_scanned = raw_rows.len();
let mut records = Vec::new();
for (point_id, content, metadata_json) in raw_rows {
let content = match content {
Some(bytes) => decode_memory_text(point_id, bytes)?,
None => return Err(MemoryStoreError::IndexDivergence { point_id }),
};
let metadata: MemoryMetadata = serde_json::from_slice(&metadata_json)
.map_err(|e| MemoryStoreError::Embedding(e.to_string()))?;
if !request.admits(&metadata) {
continue;
}
records.push(MemoryRecord { content, metadata });
}
let next_offset =
more_raw_rows_remain.then(|| request.offset.saturating_add(rows_scanned));
Ok(MemoryEnumerationPage {
records,
next_offset,
})
})
.await
.map_err(|e| MemoryStoreError::TaskJoin(format!("enumerate task join failed: {e}")))?
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use meerkat_core::memory::{
MemoryIndexBatch, MemoryIndexRequest, MemoryIndexScope, MemoryMetadata, MemorySource,
MessageRange,
};
use meerkat_core::types::SessionId;
use std::time::{SystemTime, UNIX_EPOCH};
use tempfile::TempDir;
fn meta(session_id: &SessionId) -> MemoryMetadata {
MemoryMetadata {
session_id: session_id.clone(),
source: MemorySource::Compaction {
source_range: MessageRange::single(0),
},
indexed_at: SystemTime::now(),
}
}
fn request(content: impl Into<String>, session_id: &SessionId) -> MemoryIndexRequest {
MemoryIndexRequest::new(
MemoryIndexScope::for_session(session_id.clone()),
meerkat_core::MemoryIndexableContent::Indexable(content.into()),
meta(session_id),
)
.unwrap()
}
fn request_with(
content: impl Into<String>,
session_id: &SessionId,
source_range: MessageRange,
indexed_at: SystemTime,
) -> MemoryIndexRequest {
MemoryIndexRequest::new(
MemoryIndexScope::for_session(session_id.clone()),
meerkat_core::MemoryIndexableContent::Indexable(content.into()),
MemoryMetadata {
session_id: session_id.clone(),
source: MemorySource::Compaction { source_range },
indexed_at,
},
)
.unwrap()
}
fn enumeration(limit: usize, offset: usize) -> MemoryEnumerationRequest {
MemoryEnumerationRequest {
limit,
offset,
source_overlap: None,
indexed_after: None,
}
}
fn create_pre_migration_db(db_path: &std::path::Path, rows: &[(i64, &SessionId, &str)]) {
let conn = Connection::open(db_path).unwrap();
conn.execute_batch(
"CREATE TABLE memory_metadata (
point_id INTEGER PRIMARY KEY,
metadata_json BLOB NOT NULL
);
CREATE TABLE memory_text (
point_id INTEGER PRIMARY KEY,
content BLOB NOT NULL
)",
)
.unwrap();
for (point_id, session_id, text) in rows {
let meta_json = serde_json::to_vec(&meta(session_id)).unwrap();
conn.execute(
"INSERT INTO memory_metadata (point_id, metadata_json) VALUES (?1, ?2)",
params![point_id, meta_json],
)
.unwrap();
conn.execute(
"INSERT INTO memory_text (point_id, content) VALUES (?1, ?2)",
params![point_id, text.as_bytes()],
)
.unwrap();
}
}
fn insert_old_binary_row(
db_path: &std::path::Path,
point_id: i64,
session_id: &SessionId,
text: &str,
) {
let conn = Connection::open(db_path).unwrap();
let meta_json = serde_json::to_vec(&meta(session_id)).unwrap();
conn.execute(
"INSERT INTO memory_metadata (point_id, metadata_json) VALUES (?1, ?2)",
params![point_id, meta_json],
)
.unwrap();
conn.execute(
"INSERT INTO memory_text (point_id, content) VALUES (?1, ?2)",
params![point_id, text.as_bytes()],
)
.unwrap();
}
fn query_i64(db_path: &std::path::Path, sql: &str) -> i64 {
let conn = Connection::open(db_path).unwrap();
conn.query_row(sql, [], |row| row.get(0)).unwrap()
}
#[tokio::test]
async fn test_hnsw_index_and_search() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
let other_session_id = SessionId::new();
store
.index_scoped(request(
"The user wants to implement a REST API with authentication",
&session_id,
))
.await
.unwrap();
store
.index_scoped(request(
"Configuration files use TOML format for settings",
&session_id,
))
.await
.unwrap();
store
.index_scoped(request(
"JWT tokens handle authentication and authorization",
&other_session_id,
))
.await
.unwrap();
let results = store
.search(&scope, "REST API authentication", 10)
.await
.unwrap();
assert!(!results.is_empty());
assert!(
results
.iter()
.all(|result| scope.includes(&result.metadata))
);
assert!(
results[0].content.contains("REST") || results[0].content.contains("authentication"),
"Top result should be relevant: {}",
results[0].content
);
}
#[tokio::test]
async fn test_hnsw_search_empty_store() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let scope = MemorySearchScope::for_session(SessionId::new());
let results = store.search(&scope, "anything", 10).await.unwrap();
assert!(results.is_empty());
}
#[tokio::test]
async fn test_hnsw_search_limit() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
for i in 0..10 {
store
.index_scoped(request(
format!("Item {i} with keyword test data"),
&session_id,
))
.await
.unwrap();
}
let results = store.search(&scope, "test", 3).await.unwrap();
assert!(results.len() <= 3);
}
#[tokio::test]
async fn test_hnsw_search_scopes_before_candidate_selection() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
for _ in 0..32 {
let other_session_id = SessionId::new();
store
.index_scoped(request(
"needle recall exact global candidate",
&other_session_id,
))
.await
.unwrap();
}
store
.index_scoped(request("needle recall scoped survivor", &session_id))
.await
.unwrap();
let results = store
.search(&scope, "needle recall exact global candidate", 1)
.await
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].metadata.session_id, session_id);
assert!(
results[0].content.contains("scoped survivor"),
"scoped candidates must be ranked before the limit is applied"
);
}
#[tokio::test]
async fn test_partial_hnsw_failure_repairs_live_index_no_phantom_neighbors() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("alpha survivor entry one", &session_id))
.await
.unwrap();
assert_eq!(store.hnsw_point_count(), 1);
store.arm_hnsw_insert_failure_after(1);
let batch = MemoryIndexBatch::new(
MemoryIndexScope::for_session(session_id.clone()),
vec![
request("beta doomed entry two", &session_id),
request("gamma doomed entry three", &session_id),
request("delta doomed entry four", &session_id),
],
)
.unwrap();
let err = store.index_scoped_batch(batch).await.unwrap_err();
assert_eq!(err.error_code(), "memory_lock_poisoned");
store.arm_hnsw_insert_failure_after(-1);
assert_eq!(
store.hnsw_point_count(),
1,
"repaired live index must contain exactly the surviving DB rows"
);
let results = store
.search(&scope, "beta doomed entry two", 10)
.await
.unwrap();
for result in &results {
assert!(
result.content.contains("survivor"),
"search must only return DB-present content, got: {}",
result.content
);
}
assert!(
results.iter().all(|r| scope.includes(&r.metadata)),
"all results stay within the requested scope"
);
let survivor = store
.search(&scope, "alpha survivor entry one", 1)
.await
.unwrap();
assert_eq!(survivor.len(), 1);
assert!(survivor[0].content.contains("survivor"));
}
#[tokio::test]
async fn test_hnsw_many_small_scopes_use_bounded_index_hints_across_reopen() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let mut session_ids = Vec::new();
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
for i in 0..48 {
let session_id = SessionId::new();
store
.index_scoped(request(
format!("single scoped memory entry {i}"),
&session_id,
))
.await
.unwrap();
session_ids.push(session_id);
}
assert_eq!(
store.hnsw_index_count(),
session_ids.len(),
"one-entry scopes keep separate scoped indexes for recall"
);
assert!(
store
.hnsw_index_hints()
.iter()
.all(|hint| *hint == MIN_INDEX_ELEMENTS_HINT),
"many one-entry scopes must not allocate oversized HNSW indexes"
);
assert_eq!(store.hnsw_point_count(), session_ids.len());
}
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
assert_eq!(
store.hnsw_index_count(),
0,
"open is lazy: no scoped indexes are built until first use"
);
let last_session = session_ids.last().unwrap();
let results = store
.search(
&MemorySearchScope::for_session(last_session.clone()),
"single scoped memory entry 47",
1,
)
.await
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].metadata.session_id, *last_session);
assert_eq!(
store.hnsw_index_count(),
1,
"only the searched scope is loaded"
);
assert!(
store
.hnsw_index_hints()
.iter()
.all(|hint| *hint == MIN_INDEX_ELEMENTS_HINT),
"lazily loaded one-entry scoped indexes keep bounded hints"
);
assert_eq!(store.hnsw_point_count(), 1);
}
}
#[tokio::test]
async fn test_hnsw_persists_across_reopen() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index_scoped(request(
"Persistent memory entry about Rust programming",
&session_id,
))
.await
.unwrap();
}
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let results = store.search(&scope, "Rust programming", 5).await.unwrap();
assert!(!results.is_empty(), "Data should survive reopen");
assert!(results[0].content.contains("Rust"));
}
}
#[tokio::test]
async fn test_hnsw_score_range() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("Exact match query text here", &session_id))
.await
.unwrap();
let results = store
.search(&scope, "Exact match query text here", 1)
.await
.unwrap();
assert!(!results.is_empty());
assert!(
results[0].score > 0.9,
"Exact match should have high score, got: {}",
results[0].score
);
assert!(results[0].score <= 1.0);
assert!(results[0].score >= 0.0);
}
struct ConstantEmbeddingModel {
dim: usize,
}
impl EmbeddingModel for ConstantEmbeddingModel {
fn dimension(&self) -> usize {
self.dim
}
fn embed(&self, _text: &str) -> Vec<f32> {
let mut v = vec![0.0f32; self.dim];
v[0] = 1.0;
v
}
}
#[tokio::test]
async fn test_injected_policy_is_ranking_authority() {
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
let default_score = {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
store
.index_scoped(request("alpha beta gamma", &session_id))
.await
.unwrap();
let results = store
.search(&scope, "completely unrelated query", 1)
.await
.unwrap();
results.first().map(|r| r.score)
};
let constant_score = {
let dir = TempDir::new().unwrap();
let policy = MemoryRankingPolicy::new(
Arc::new(ConstantEmbeddingModel { dim: 16 }),
HnswParams::default(),
);
let store =
HnswMemoryStore::open_with_policy(dir.path().join("memory"), policy).unwrap();
store
.index_scoped(request("alpha beta gamma", &session_id))
.await
.unwrap();
let results = store
.search(&scope, "completely unrelated query", 1)
.await
.unwrap();
results.first().map(|r| r.score)
};
let constant_score = constant_score.expect("constant policy matches all content");
assert!(
constant_score > 0.99,
"constant embedding policy must rank unrelated text as a match, got {constant_score}"
);
assert!(
default_score.map(|s| s < 0.99).unwrap_or(true),
"default policy must not rank unrelated text as a perfect match"
);
}
fn corrupt_text_bytes(db_path: &std::path::Path, bytes: &[u8]) {
let conn = Connection::open(db_path).unwrap();
let updated = conn
.execute("UPDATE memory_text SET content = ?1", params![bytes])
.unwrap();
assert!(updated > 0, "corruption fixture must hit at least one row");
}
#[tokio::test]
async fn test_corrupt_text_bytes_fail_closed_on_first_scope_load() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let session_id = SessionId::new();
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index_scoped(request("clean entry before corruption", &session_id))
.await
.unwrap();
}
corrupt_text_bytes(&memory_dir.join("memory.sqlite3"), &[0xff, 0xfe, 0x41]);
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let scope = MemorySearchScope::for_session(session_id);
let err = store
.search(&scope, "clean entry before corruption", 5)
.await
.unwrap_err();
assert_eq!(err.error_code(), "memory_text_corruption");
}
#[tokio::test]
async fn test_corrupt_text_bytes_fail_closed_on_search() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("entry destined for corruption", &session_id))
.await
.unwrap();
corrupt_text_bytes(&memory_dir.join("memory.sqlite3"), &[0xc3, 0x28]);
let err = store
.search(&scope, "entry destined for corruption", 5)
.await
.unwrap_err();
assert_eq!(err.error_code(), "memory_text_corruption");
}
#[tokio::test]
async fn test_missing_durable_row_is_typed_divergence_not_skip() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("row deleted out of band", &session_id))
.await
.unwrap();
{
let conn = Connection::open(memory_dir.join("memory.sqlite3")).unwrap();
conn.execute("DELETE FROM memory_text", []).unwrap();
}
let err = store
.search(&scope, "row deleted out of band", 5)
.await
.unwrap_err();
assert_eq!(err.error_code(), "memory_index_divergence");
}
#[tokio::test]
async fn test_repair_failure_poisons_scope_then_next_index_self_heals() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let db_path = memory_dir.join("memory.sqlite3");
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("survivor pending corruption", &session_id))
.await
.unwrap();
corrupt_text_bytes(&db_path, &[0xff, 0x00, 0x41]);
store.arm_hnsw_insert_failure_after(1);
let batch = MemoryIndexBatch::new(
MemoryIndexScope::for_session(session_id.clone()),
vec![
request("doomed one", &session_id),
request("doomed two", &session_id),
],
)
.unwrap();
let err = store.index_scoped_batch(batch).await.unwrap_err();
assert_eq!(err.error_code(), "memory_scope_repair_failed");
store.arm_hnsw_insert_failure_after(-1);
let err = store.search(&scope, "survivor", 5).await.unwrap_err();
assert_eq!(err.error_code(), "memory_scope_poisoned");
{
let conn = Connection::open(&db_path).unwrap();
conn.execute(
"UPDATE memory_text SET content = ?1",
params![b"survivor restored text".as_slice()],
)
.unwrap();
}
store
.index_scoped(request("fresh entry after heal", &session_id))
.await
.unwrap();
let results = store
.search(&scope, "fresh entry after heal", 5)
.await
.unwrap();
assert!(
results
.iter()
.any(|r| r.content.contains("fresh entry after heal")),
"self-healed scope must serve durable-derived candidates"
);
let survivor = store
.search(&scope, "survivor restored text", 5)
.await
.unwrap();
assert!(
survivor
.iter()
.any(|r| r.content.contains("survivor restored text")),
"self-healed scope must include pre-existing durable rows"
);
}
#[tokio::test]
async fn test_distinct_failures_surface_as_distinct_typed_variants() {
assert_eq!(
MemoryStoreError::Embedding("x".into()).error_code(),
"memory_embedding"
);
assert_eq!(
MemoryStoreError::Storage("x".into()).error_code(),
"memory_storage"
);
assert_eq!(
MemoryStoreError::LockPoisoned.error_code(),
"memory_lock_poisoned"
);
assert_eq!(
MemoryStoreError::PointIdOverflow.error_code(),
"memory_point_id_overflow"
);
assert_eq!(
MemoryStoreError::PointIdOutOfRange.error_code(),
"memory_point_id_out_of_range"
);
assert_eq!(
MemoryStoreError::TextCorruption { point_id: 7 }.error_code(),
"memory_text_corruption"
);
assert_eq!(
MemoryStoreError::IndexDivergence { point_id: 7 }.error_code(),
"memory_index_divergence"
);
assert_eq!(
MemoryStoreError::ScopePoisoned.error_code(),
"memory_scope_poisoned"
);
assert_eq!(
MemoryStoreError::ScopeRepairFailed {
original: Box::new(MemoryStoreError::LockPoisoned),
repair: Box::new(MemoryStoreError::TextCorruption { point_id: 7 }),
}
.error_code(),
"memory_scope_repair_failed"
);
}
#[tokio::test]
async fn test_migration_from_pre_session_id_schema_heals_and_serves() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
std::fs::create_dir_all(&memory_dir).unwrap();
let db_path = memory_dir.join("memory.sqlite3");
let session_a = SessionId::new();
let session_b = SessionId::new();
create_pre_migration_db(
&db_path,
&[
(0, &session_a, "alpha entry from the old world"),
(1, &session_a, "beta entry from the old world"),
(2, &session_b, "gamma entry in another scope"),
],
);
let store = HnswMemoryStore::open(&memory_dir).unwrap();
assert_eq!(
query_i64(
&db_path,
"SELECT COUNT(*) FROM memory_metadata WHERE session_id IS NULL",
),
0
);
assert_eq!(
query_i64(
&db_path,
"SELECT high_water FROM memory_allocator WHERE id = 0"
),
3
);
let scope_a = MemorySearchScope::for_session(session_a.clone());
let results = store
.search(&scope_a, "alpha entry from the old world", 10)
.await
.unwrap();
assert!(!results.is_empty());
assert!(results.iter().all(|r| r.metadata.session_id == session_a));
let page = store
.enumerate_scoped(&scope_a, enumeration(10, 0))
.await
.unwrap();
assert_eq!(page.records.len(), 2);
assert!(page.records[0].content.contains("alpha"));
assert!(page.records[1].content.contains("beta"));
assert_eq!(page.next_offset, None);
store
.index_scoped(request("delta entry post migration", &session_a))
.await
.unwrap();
assert_eq!(
query_i64(&db_path, "SELECT MAX(point_id) FROM memory_metadata"),
3
);
}
#[tokio::test]
async fn test_migration_is_idempotent_across_reopens() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let session_id = SessionId::new();
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index_scoped(request("entry surviving reopens", &session_id))
.await
.unwrap();
}
{
let _store = HnswMemoryStore::open(&memory_dir).unwrap();
}
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let db_path = memory_dir.join("memory.sqlite3");
let conn = Connection::open(&db_path).unwrap();
let session_id_columns: i64 = conn
.query_row(
"SELECT COUNT(*) FROM pragma_table_info('memory_metadata') WHERE name = 'session_id'",
[],
|row| row.get(0),
)
.unwrap();
assert_eq!(
session_id_columns, 1,
"repeated opens must not re-add the column"
);
assert_eq!(
query_i64(&db_path, "SELECT COUNT(*) FROM memory_allocator"),
1,
"allocator stays a single row"
);
let scope = MemorySearchScope::for_session(session_id);
let results = store
.search(&scope, "entry surviving reopens", 5)
.await
.unwrap();
assert_eq!(results.len(), 1);
}
#[tokio::test]
async fn test_null_session_id_row_healed_by_lazy_search_load() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let db_path = memory_dir.join("memory.sqlite3");
let session_id = SessionId::new();
let store = HnswMemoryStore::open(&memory_dir).unwrap();
insert_old_binary_row(&db_path, 100, &session_id, "row from an old binary");
let scope = MemorySearchScope::for_session(session_id.clone());
let results = store
.search(&scope, "row from an old binary", 5)
.await
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].metadata.session_id, session_id);
assert_eq!(
query_i64(
&db_path,
"SELECT COUNT(*) FROM memory_metadata WHERE session_id IS NULL",
),
0,
"the lazy load healed the projection cell"
);
}
#[tokio::test]
async fn test_null_session_id_row_healed_before_enumerate_and_drop() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let db_path = memory_dir.join("memory.sqlite3");
let session_id = SessionId::new();
let store = HnswMemoryStore::open(&memory_dir).unwrap();
insert_old_binary_row(&db_path, 100, &session_id, "old binary row to enumerate");
let scope = MemorySearchScope::for_session(session_id.clone());
let page = store
.enumerate_scoped(&scope, enumeration(10, 0))
.await
.unwrap();
assert_eq!(page.records.len(), 1);
assert!(page.records[0].content.contains("old binary row"));
insert_old_binary_row(&db_path, 101, &session_id, "second old binary row");
let receipt = store
.drop_scope(&MemoryOwner::canonical_session(session_id.clone()))
.await
.unwrap();
assert_eq!(receipt.dropped_entries, 2);
assert_eq!(receipt.owner.session_id(), &session_id);
assert_eq!(
query_i64(&db_path, "SELECT COUNT(*) FROM memory_metadata"),
0
);
assert_eq!(query_i64(&db_path, "SELECT COUNT(*) FROM memory_text"), 0);
}
#[tokio::test]
async fn test_old_binary_row_does_not_wedge_point_allocation() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let db_path = memory_dir.join("memory.sqlite3");
let session_id = SessionId::new();
let store = HnswMemoryStore::open(&memory_dir).unwrap();
insert_old_binary_row(&db_path, 100, &session_id, "old binary high row");
store
.index_scoped(request("new binary row after the gap", &session_id))
.await
.unwrap();
assert_eq!(
query_i64(&db_path, "SELECT MAX(point_id) FROM memory_metadata"),
101,
"allocation must jump past the old-binary row"
);
assert_eq!(
query_i64(
&db_path,
"SELECT high_water FROM memory_allocator WHERE id = 0"
),
102
);
}
#[tokio::test]
async fn test_index_into_unloaded_existing_scope_serves_prior_entries() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index_scoped(request("prior alpha entry", &session_id))
.await
.unwrap();
}
let store = HnswMemoryStore::open(&memory_dir).unwrap();
assert_eq!(store.hnsw_index_count(), 0, "open must not load scopes");
store
.index_scoped(request("fresh beta entry", &session_id))
.await
.unwrap();
let prior = store.search(&scope, "prior alpha entry", 5).await.unwrap();
assert!(
prior
.iter()
.any(|r| r.content.contains("prior alpha entry")),
"prior durable entries must stay searchable after an insert into an unloaded scope"
);
let fresh = store.search(&scope, "fresh beta entry", 5).await.unwrap();
assert!(fresh.iter().any(|r| r.content.contains("fresh beta entry")));
assert_eq!(store.hnsw_point_count(), 2);
}
#[tokio::test]
async fn test_drop_scope_removes_rows_durably_and_preserves_other_scopes() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let db_path = memory_dir.join("memory.sqlite3");
let session_a = SessionId::new();
let session_b = SessionId::new();
let scope_a = MemorySearchScope::for_session(session_a.clone());
let scope_b = MemorySearchScope::for_session(session_b.clone());
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index_scoped(request("doomed alpha entry", &session_a))
.await
.unwrap();
store
.index_scoped(request("doomed beta entry", &session_a))
.await
.unwrap();
store
.index_scoped(request("surviving gamma entry", &session_b))
.await
.unwrap();
let receipt = store
.drop_scope(&MemoryOwner::canonical_session(session_a.clone()))
.await
.unwrap();
assert_eq!(receipt.dropped_entries, 2);
assert_eq!(receipt.owner.session_id(), &session_a);
let dropped = store
.search(&scope_a, "doomed alpha entry", 10)
.await
.unwrap();
assert!(dropped.is_empty(), "dropped scope must serve nothing");
let page = store
.enumerate_scoped(&scope_a, enumeration(10, 0))
.await
.unwrap();
assert!(page.records.is_empty());
assert_eq!(page.next_offset, None);
let surviving = store
.search(&scope_b, "surviving gamma entry", 10)
.await
.unwrap();
assert_eq!(surviving.len(), 1);
}
assert_eq!(
query_i64(&db_path, "SELECT COUNT(*) FROM memory_metadata"),
1
);
assert_eq!(query_i64(&db_path, "SELECT COUNT(*) FROM memory_text"), 1);
let store = HnswMemoryStore::open(&memory_dir).unwrap();
assert!(
store
.search(&scope_a, "doomed alpha entry", 10)
.await
.unwrap()
.is_empty()
);
assert_eq!(
store
.search(&scope_b, "surviving gamma entry", 10)
.await
.unwrap()
.len(),
1
);
}
#[tokio::test]
async fn test_drop_scope_unknown_scope_reports_zero() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let receipt = store
.drop_scope(&MemoryOwner::canonical_session(SessionId::new()))
.await
.unwrap();
assert_eq!(receipt.dropped_entries, 0);
}
#[tokio::test]
async fn test_point_ids_never_reused_after_drop_scope() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let db_path = memory_dir.join("memory.sqlite3");
let session_id = SessionId::new();
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index_scoped(request("first doomed entry", &session_id))
.await
.unwrap();
store
.index_scoped(request("second doomed entry", &session_id))
.await
.unwrap();
assert_eq!(
query_i64(&db_path, "SELECT MAX(point_id) FROM memory_metadata"),
1
);
store
.drop_scope(&MemoryOwner::canonical_session(session_id.clone()))
.await
.unwrap();
store
.index_scoped(request("entry after the drop", &session_id))
.await
.unwrap();
assert_eq!(
query_i64(&db_path, "SELECT MAX(point_id) FROM memory_metadata"),
2,
"IDs must strictly increase across a drop, never recycle"
);
assert_eq!(
query_i64(
&db_path,
"SELECT high_water FROM memory_allocator WHERE id = 0"
),
3
);
}
#[tokio::test]
async fn test_enumerate_scoped_pages_deterministically_in_insertion_order() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let other_session = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
let texts = [
"entry number zero",
"entry number one",
"entry number two",
"entry number three",
"entry number four",
];
for (i, text) in texts.iter().enumerate() {
store
.index_scoped(request(*text, &session_id))
.await
.unwrap();
store
.index_scoped(request(format!("interloper {i}"), &other_session))
.await
.unwrap();
}
let first = store
.enumerate_scoped(&scope, enumeration(2, 0))
.await
.unwrap();
assert_eq!(first.records.len(), 2);
assert_eq!(first.records[0].content, "entry number zero");
assert_eq!(first.records[1].content, "entry number one");
assert_eq!(first.next_offset, Some(2));
let second = store
.enumerate_scoped(&scope, enumeration(2, 2))
.await
.unwrap();
assert_eq!(second.records.len(), 2);
assert_eq!(second.records[0].content, "entry number two");
assert_eq!(second.records[1].content, "entry number three");
assert_eq!(second.next_offset, Some(4));
let last = store
.enumerate_scoped(&scope, enumeration(2, 4))
.await
.unwrap();
assert_eq!(last.records.len(), 1);
assert_eq!(last.records[0].content, "entry number four");
assert_eq!(last.next_offset, None);
let replay = store
.enumerate_scoped(&scope, enumeration(2, 0))
.await
.unwrap();
assert_eq!(replay.records.len(), 2);
assert_eq!(replay.records[0].content, "entry number zero");
assert_eq!(replay.records[1].content, "entry number one");
assert_eq!(replay.next_offset, Some(2));
}
#[tokio::test]
async fn test_enumerate_scoped_source_overlap_filters_post_deserialize() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
let indexed_at = UNIX_EPOCH + Duration::from_secs(1_000);
store
.index_scoped(request_with(
"covers zero to five",
&session_id,
MessageRange::new(0, 5).unwrap(),
indexed_at,
))
.await
.unwrap();
store
.index_scoped(request_with(
"covers five to ten",
&session_id,
MessageRange::new(5, 10).unwrap(),
indexed_at,
))
.await
.unwrap();
store
.index_scoped(request_with(
"covers ten to fifteen",
&session_id,
MessageRange::new(10, 15).unwrap(),
indexed_at,
))
.await
.unwrap();
let page = store
.enumerate_scoped(
&scope,
MemoryEnumerationRequest {
limit: 10,
offset: 0,
source_overlap: Some(MessageRange::new(4, 6).unwrap()),
indexed_after: None,
},
)
.await
.unwrap();
assert_eq!(page.records.len(), 2);
assert_eq!(page.records[0].content, "covers zero to five");
assert_eq!(page.records[1].content, "covers five to ten");
assert_eq!(page.next_offset, None);
let narrow = store
.enumerate_scoped(
&scope,
MemoryEnumerationRequest {
limit: 10,
offset: 0,
source_overlap: Some(MessageRange::new(0, 1).unwrap()),
indexed_after: None,
},
)
.await
.unwrap();
assert_eq!(narrow.records.len(), 1);
assert_eq!(narrow.records[0].content, "covers zero to five");
assert_eq!(narrow.next_offset, None);
}
#[tokio::test]
async fn test_enumerate_scoped_indexed_after_filter() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
let generation_one = UNIX_EPOCH + Duration::from_secs(1_000);
let generation_two = UNIX_EPOCH + Duration::from_secs(2_000);
store
.index_scoped(request_with(
"first generation summary",
&session_id,
MessageRange::new(0, 3).unwrap(),
generation_one,
))
.await
.unwrap();
store
.index_scoped(request_with(
"second generation summary",
&session_id,
MessageRange::new(0, 3).unwrap(),
generation_two,
))
.await
.unwrap();
let page = store
.enumerate_scoped(
&scope,
MemoryEnumerationRequest {
limit: 10,
offset: 0,
source_overlap: None,
indexed_after: Some(generation_one),
},
)
.await
.unwrap();
assert_eq!(page.records.len(), 1);
assert_eq!(page.records[0].content, "second generation summary");
let none_left = store
.enumerate_scoped(
&scope,
MemoryEnumerationRequest {
limit: 10,
offset: 0,
source_overlap: None,
indexed_after: Some(generation_two),
},
)
.await
.unwrap();
assert!(none_left.records.is_empty());
}
#[tokio::test]
async fn test_enumerate_scoped_corruption_fails_closed() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("entry destined for corruption", &session_id))
.await
.unwrap();
corrupt_text_bytes(&memory_dir.join("memory.sqlite3"), &[0xc3, 0x28]);
let err = store
.enumerate_scoped(&scope, enumeration(10, 0))
.await
.unwrap_err();
assert_eq!(err.error_code(), "memory_text_corruption");
}
#[tokio::test]
async fn test_enumerate_scoped_missing_text_row_is_typed_divergence() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("row deleted out of band", &session_id))
.await
.unwrap();
{
let conn = Connection::open(memory_dir.join("memory.sqlite3")).unwrap();
conn.execute("DELETE FROM memory_text", []).unwrap();
}
let err = store
.enumerate_scoped(&scope, enumeration(10, 0))
.await
.unwrap_err();
assert_eq!(err.error_code(), "memory_index_divergence");
}
#[tokio::test]
async fn test_enumerate_scoped_empty_scope() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let scope = MemorySearchScope::for_session(SessionId::new());
let page = store
.enumerate_scoped(&scope, enumeration(10, 0))
.await
.unwrap();
assert!(page.records.is_empty());
assert_eq!(page.next_offset, None);
}
#[tokio::test]
async fn test_enumerate_scoped_limit_zero_is_typed_error() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
let session_id = SessionId::new();
let scope = MemorySearchScope::for_session(session_id.clone());
store
.index_scoped(request("entry".to_string(), &session_id))
.await
.unwrap();
let error = store
.enumerate_scoped(&scope, enumeration(0, 1))
.await
.expect_err("limit zero must be rejected");
assert!(matches!(error, MemoryStoreError::EnumerationLimitZero));
assert_eq!(error.error_code(), "memory_enumeration_limit_zero");
}
}