use async_trait::async_trait;
use paladin_core::platform::container::sanctum::SanctumEntry;
use paladin_ports::output::sanctum_port::{
SanctumError, SanctumFilter, SanctumPort, SanctumQuery, SanctumSearchResult,
};
use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, RwLock};
use uuid::Uuid;
#[derive(Debug, Clone)]
pub struct InMemorySanctumConfig {
pub max_entries: usize,
}
impl Default for InMemorySanctumConfig {
fn default() -> Self {
Self {
max_entries: 10_000,
}
}
}
pub struct InMemorySanctum {
storage: Arc<RwLock<HashMap<Uuid, SanctumEntry>>>,
lru_queue: Arc<RwLock<VecDeque<Uuid>>>,
config: InMemorySanctumConfig,
}
impl InMemorySanctum {
pub fn new(max_entries: usize) -> Self {
Self {
storage: Arc::new(RwLock::new(HashMap::new())),
lru_queue: Arc::new(RwLock::new(VecDeque::new())),
config: InMemorySanctumConfig { max_entries },
}
}
pub fn with_config(config: InMemorySanctumConfig) -> Self {
Self {
storage: Arc::new(RwLock::new(HashMap::new())),
lru_queue: Arc::new(RwLock::new(VecDeque::new())),
config,
}
}
fn cosine_similarity(a: &[f32], b: &[f32]) -> Result<f32, SanctumError> {
if a.len() != b.len() {
return Err(SanctumError::InvalidDimension(format!(
"Vector dimensions don't match: {} vs {}",
a.len(),
b.len()
)));
}
if a.is_empty() {
return Err(SanctumError::InvalidDimension(
"Cannot calculate similarity of empty vectors".to_string(),
));
}
let dot_product: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
let mag_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let mag_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if mag_a == 0.0 || mag_b == 0.0 {
return Ok(0.0); }
Ok(dot_product / (mag_a * mag_b))
}
fn matches_filter(entry: &SanctumEntry, filter: &SanctumFilter) -> bool {
if let Some(ref paladin_id) = filter.paladin_id
&& &entry.memory.paladin_id != paladin_id
{
return false;
}
if let Some(memory_type) = filter.memory_type
&& entry.memory.memory_type != memory_type
{
return false;
}
if let Some(created_after) = filter.created_after
&& entry.memory.created_at < created_after
{
return false;
}
if let Some(created_before) = filter.created_before
&& entry.memory.created_at > created_before
{
return false;
}
if let Some(min_importance) = filter.min_importance
&& entry.memory.importance < min_importance
{
return false;
}
for (key, value) in &filter.metadata_filters {
match entry.memory.metadata.get(key) {
Some(entry_value) if entry_value == value => {}
_ => return false,
}
}
true
}
fn evict_if_needed(&self) {
let storage = self
.storage
.read()
.expect("Failed to acquire read lock on storage");
if storage.len() >= self.config.max_entries {
drop(storage);
let mut lru = self
.lru_queue
.write()
.expect("Failed to acquire write lock on LRU queue");
let mut storage = self
.storage
.write()
.expect("Failed to acquire write lock on storage");
if let Some(oldest_id) = lru.pop_front() {
storage.remove(&oldest_id);
}
}
}
fn touch_entry(&self, id: &Uuid) {
let mut lru = self
.lru_queue
.write()
.expect("Failed to acquire write lock on LRU queue");
if let Some(pos) = lru.iter().position(|x| x == id) {
lru.remove(pos);
}
lru.push_back(*id);
}
}
#[async_trait]
impl SanctumPort for InMemorySanctum {
async fn store(&self, entry: SanctumEntry) -> Result<(), SanctumError> {
self.evict_if_needed();
let id = entry.memory.id;
let mut storage = self
.storage
.write()
.expect("Failed to acquire write lock on storage");
storage.insert(id, entry);
drop(storage);
self.touch_entry(&id);
Ok(())
}
async fn store_batch(&self, entries: Vec<SanctumEntry>) -> Result<(), SanctumError> {
for entry in entries {
self.store(entry).await?;
}
Ok(())
}
async fn search(&self, query: SanctumQuery) -> Result<Vec<SanctumSearchResult>, SanctumError> {
let storage = self
.storage
.read()
.expect("Failed to acquire read lock on storage");
let mut results = Vec::new();
for entry in storage.values() {
if let Some(ref filter) = query.filter
&& !Self::matches_filter(entry, filter)
{
continue;
}
let score = Self::cosine_similarity(&query.embedding, &entry.embedding)?;
if let Some(min_score) = query.min_score
&& score < min_score
{
continue;
}
results.push(SanctumSearchResult {
entry: entry.clone(),
score,
});
}
results.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
results.truncate(query.top_k);
Ok(results)
}
async fn delete(&self, id: &str) -> Result<bool, SanctumError> {
let uuid = Uuid::parse_str(id)
.map_err(|e| SanctumError::NotFound(format!("Invalid UUID format: {}", e)))?;
let mut storage = self
.storage
.write()
.expect("Failed to acquire write lock on storage");
let removed = storage.remove(&uuid).is_some();
if removed {
drop(storage);
let mut lru = self
.lru_queue
.write()
.expect("Failed to acquire write lock on LRU queue");
if let Some(pos) = lru.iter().position(|x| x == &uuid) {
lru.remove(pos);
}
}
Ok(removed)
}
async fn update(&self, entry: SanctumEntry) -> Result<(), SanctumError> {
let id = entry.memory.id;
let mut storage = self
.storage
.write()
.expect("Failed to acquire write lock on storage");
if !storage.contains_key(&id) {
return Err(SanctumError::NotFound(format!("Entry not found: {}", id)));
}
storage.insert(id, entry);
drop(storage);
self.touch_entry(&id);
Ok(())
}
async fn count(&self, filter: Option<SanctumFilter>) -> Result<usize, SanctumError> {
let storage = self
.storage
.read()
.expect("Failed to acquire read lock on storage");
if let Some(filter) = filter {
Ok(storage
.values()
.filter(|entry| Self::matches_filter(entry, &filter))
.count())
} else {
Ok(storage.len())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cosine_similarity_identical() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
let similarity = InMemorySanctum::cosine_similarity(&a, &b).unwrap();
assert!((similarity - 1.0).abs() < 1e-6);
}
#[test]
fn test_cosine_similarity_orthogonal() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![0.0, 1.0, 0.0];
let similarity = InMemorySanctum::cosine_similarity(&a, &b).unwrap();
assert!((similarity - 0.0).abs() < 1e-6);
}
#[test]
fn test_cosine_similarity_opposite() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![-1.0, 0.0, 0.0];
let similarity = InMemorySanctum::cosine_similarity(&a, &b).unwrap();
assert!((similarity - (-1.0)).abs() < 1e-6);
}
#[test]
fn test_cosine_similarity_dimension_mismatch() {
let a = vec![1.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
let result = InMemorySanctum::cosine_similarity(&a, &b);
assert!(result.is_err());
}
#[test]
fn test_cosine_similarity_empty() {
let a: Vec<f32> = vec![];
let b: Vec<f32> = vec![];
let result = InMemorySanctum::cosine_similarity(&a, &b);
assert!(result.is_err());
}
#[test]
fn test_cosine_similarity_zero_vector() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
let similarity = InMemorySanctum::cosine_similarity(&a, &b).unwrap();
assert_eq!(similarity, 0.0);
}
}