use async_trait::async_trait;
use hnsw_rs::prelude::{DistCosine, Hnsw};
use meerkat_core::memory::{MemoryMetadata, MemoryResult, MemoryStore, MemoryStoreError};
use redb::{Database, ReadableTable, TableDefinition};
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Mutex;
const METADATA_TABLE: TableDefinition<u64, &[u8]> = TableDefinition::new("memory_metadata");
const TEXT_TABLE: TableDefinition<u64, &[u8]> = TableDefinition::new("memory_text");
const VOCAB_DIM: usize = 4096;
const MAX_NB_CONNECTION: usize = 16;
const MAX_LAYER: usize = 16;
const EF_CONSTRUCTION: usize = 200;
const DEFAULT_MAX_ELEMENTS: usize = 100_000;
pub struct HnswMemoryStore {
index: Arc<std::sync::RwLock<Hnsw<'static, f32, DistCosine>>>,
db: Arc<Database>,
next_id: AtomicUsize,
insert_lock: Mutex<()>,
path: PathBuf,
}
impl HnswMemoryStore {
pub fn open(dir: impl AsRef<Path>) -> Result<Self, MemoryStoreError> {
let dir = dir.as_ref();
std::fs::create_dir_all(dir).map_err(MemoryStoreError::Io)?;
let db_path = dir.join("memory.redb");
let db = Arc::new(
Database::create(&db_path).map_err(|e| MemoryStoreError::Index(e.to_string()))?,
);
let write_txn = db
.begin_write()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
{
let _ = write_txn
.open_table(METADATA_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let _ = write_txn
.open_table(TEXT_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
}
write_txn
.commit()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let next_id = {
let read_txn = db
.begin_read()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let table = read_txn
.open_table(METADATA_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let mut max_id = 0usize;
let iter = table
.iter()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
for entry in iter {
let (key, _) = entry.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let id = usize::try_from(key.value())
.map_err(|_| MemoryStoreError::Index("point ID out of range".to_string()))?;
if id >= max_id {
max_id = id
.checked_add(1)
.ok_or_else(|| MemoryStoreError::Index("point ID overflow".to_string()))?;
}
}
max_id
};
let hnsw = Hnsw::<'static, f32, DistCosine>::new(
MAX_NB_CONNECTION,
DEFAULT_MAX_ELEMENTS,
MAX_LAYER,
EF_CONSTRUCTION,
DistCosine {},
);
if next_id > 0 {
let read_txn = db
.begin_read()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let text_table = read_txn
.open_table(TEXT_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
for point_id in 0..next_id {
let point_id_u64 = u64::try_from(point_id)
.map_err(|_| MemoryStoreError::Index("point ID out of range".to_string()))?;
if let Some(text_guard) = text_table
.get(point_id_u64)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?
{
let text = String::from_utf8_lossy(text_guard.value());
let embedding = text_to_embedding(&text);
hnsw.insert((&embedding, point_id));
}
}
}
Ok(Self {
index: Arc::new(std::sync::RwLock::new(hnsw)),
db,
next_id: AtomicUsize::new(next_id),
insert_lock: Mutex::new(()),
path: dir.to_path_buf(),
})
}
pub fn path(&self) -> &Path {
&self.path
}
}
#[async_trait]
impl MemoryStore for HnswMemoryStore {
async fn index(&self, content: &str, metadata: MemoryMetadata) -> Result<(), MemoryStoreError> {
let meta_json = serde_json::to_vec(&metadata)
.map_err(|e| MemoryStoreError::Embedding(e.to_string()))?;
let content = content.to_owned();
let db = Arc::clone(&self.db);
let index = Arc::clone(&self.index);
let _guard = self.insert_lock.lock().await;
let point_id = self.next_id.load(Ordering::Acquire);
let point_id_u64 = u64::try_from(point_id)
.map_err(|_| MemoryStoreError::Index("point ID out of range".to_string()))?;
tokio::task::spawn_blocking(move || {
let embedding = text_to_embedding(&content);
let write_txn = db
.begin_write()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
{
let mut meta_table = write_txn
.open_table(METADATA_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let mut text_table = write_txn
.open_table(TEXT_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
meta_table
.insert(point_id_u64, meta_json.as_slice())
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
text_table
.insert(point_id_u64, content.as_bytes())
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
}
write_txn
.commit()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let index = index
.write()
.map_err(|_| MemoryStoreError::Index("HNSW index lock poisoned".to_string()))?;
index.insert((&embedding, point_id));
Ok::<(), MemoryStoreError>(())
})
.await
.map_err(|e| MemoryStoreError::Index(format!("index task join failed: {e}")))??;
let next_id = point_id
.checked_add(1)
.ok_or_else(|| MemoryStoreError::Index("point ID overflow".to_string()))?;
self.next_id.store(next_id, Ordering::Release);
Ok(())
}
async fn search(
&self,
query: &str,
limit: usize,
) -> Result<Vec<MemoryResult>, MemoryStoreError> {
if limit == 0 {
return Ok(Vec::new());
}
let query = query.to_owned();
let db = Arc::clone(&self.db);
let index = Arc::clone(&self.index);
tokio::task::spawn_blocking(move || {
let embedding = text_to_embedding(&query);
let ef_search = limit.max(EF_CONSTRUCTION);
let neighbors = {
let index = index
.read()
.map_err(|_| MemoryStoreError::Index("HNSW index lock poisoned".to_string()))?;
index.search(&embedding, limit, ef_search)
};
let read_txn = db
.begin_read()
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let meta_table = read_txn
.open_table(METADATA_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let text_table = read_txn
.open_table(TEXT_TABLE)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?;
let mut results = Vec::with_capacity(neighbors.len());
for neighbor in &neighbors {
let point_id = u64::try_from(neighbor.d_id)
.map_err(|_| MemoryStoreError::Index("point ID out of range".to_string()))?;
let content = match text_table
.get(point_id)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?
{
Some(guard) => String::from_utf8_lossy(guard.value()).into_owned(),
None => continue,
};
let metadata = match meta_table
.get(point_id)
.map_err(|e| MemoryStoreError::Index(e.to_string()))?
{
Some(guard) => serde_json::from_slice(guard.value())
.map_err(|e| MemoryStoreError::Embedding(e.to_string()))?,
None => continue,
};
let score = 1.0 - (neighbor.distance / 2.0);
results.push(MemoryResult {
content,
metadata,
score,
});
}
Ok::<Vec<MemoryResult>, MemoryStoreError>(results)
})
.await
.map_err(|e| MemoryStoreError::Index(format!("search task join failed: {e}")))?
}
}
fn text_to_embedding(text: &str) -> [f32; VOCAB_DIM] {
let mut vec = [0.0f32; 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)
}) % 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
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
use meerkat_core::types::SessionId;
use std::time::SystemTime;
use tempfile::TempDir;
fn meta() -> MemoryMetadata {
MemoryMetadata {
session_id: SessionId::new(),
turn: Some(1),
indexed_at: SystemTime::now(),
}
}
#[tokio::test]
async fn test_hnsw_index_and_search() {
let dir = TempDir::new().unwrap();
let store = HnswMemoryStore::open(dir.path().join("memory")).unwrap();
store
.index(
"The user wants to implement a REST API with authentication",
meta(),
)
.await
.unwrap();
store
.index("Configuration files use TOML format for settings", meta())
.await
.unwrap();
store
.index("JWT tokens handle authentication and authorization", meta())
.await
.unwrap();
let results = store.search("REST API authentication", 10).await.unwrap();
assert!(!results.is_empty());
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 results = store.search("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();
for i in 0..10 {
store
.index(&format!("Item {} with keyword test data", i), meta())
.await
.unwrap();
}
let results = store.search("test", 3).await.unwrap();
assert!(results.len() <= 3);
}
#[tokio::test]
async fn test_hnsw_persists_across_reopen() {
let dir = TempDir::new().unwrap();
let memory_dir = dir.path().join("memory");
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
store
.index("Persistent memory entry about Rust programming", meta())
.await
.unwrap();
}
{
let store = HnswMemoryStore::open(&memory_dir).unwrap();
let results = store.search("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();
store
.index("Exact match query text here", meta())
.await
.unwrap();
let results = store
.search("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
);
}
}