use async_trait::async_trait;
use paladin_core::platform::container::sanctum::{Memory, MemoryType, SanctumEntry};
use paladin_ports::output::sanctum_port::{
SanctumError, SanctumFilter, SanctumPort, SanctumQuery, SanctumSearchResult,
};
use qdrant_client::Qdrant;
use qdrant_client::qdrant::r#match::MatchValue;
use qdrant_client::qdrant::vectors_config::Config;
use qdrant_client::qdrant::{
Condition, Distance, Filter, PointStruct, Range, Value as QdrantValue, VectorParams,
VectorsConfig,
};
use serde_json::Value;
use std::collections::HashMap;
use uuid::Uuid;
pub struct QdrantSanctumAdapter {
client: Qdrant,
collection_name: String,
vector_dimension: usize,
}
impl QdrantSanctumAdapter {
pub async fn new(
url: &str,
collection_name: &str,
vector_dimension: usize,
) -> Result<Self, SanctumError> {
let client = Qdrant::from_url(url).build().map_err(|e| {
SanctumError::ConfigError(format!("Failed to create Qdrant client: {}", e))
})?;
let adapter = Self {
client,
collection_name: collection_name.to_string(),
vector_dimension,
};
adapter.ensure_collection_exists().await?;
Ok(adapter)
}
async fn ensure_collection_exists(&self) -> Result<(), SanctumError> {
let collections = self.client.list_collections().await.map_err(|e| {
SanctumError::StorageError(format!("Failed to list collections: {}", e))
})?;
let exists = collections
.collections
.iter()
.any(|c| c.name == self.collection_name);
if !exists {
self.client
.create_collection(
qdrant_client::qdrant::CreateCollectionBuilder::new(&self.collection_name)
.vectors_config(VectorsConfig {
config: Some(Config::Params(VectorParams {
size: self.vector_dimension as u64,
distance: Distance::Cosine.into(),
hnsw_config: None,
quantization_config: None,
on_disk: None,
datatype: None,
multivector_config: None,
})),
}),
)
.await
.map_err(|e| {
SanctumError::StorageError(format!("Failed to create collection: {}", e))
})?;
}
Ok(())
}
fn entry_to_point(&self, entry: &SanctumEntry) -> Result<PointStruct, SanctumError> {
if entry.embedding.len() != self.vector_dimension {
return Err(SanctumError::InvalidDimension(format!(
"Expected {} dimensions, got {}",
self.vector_dimension,
entry.embedding.len()
)));
}
let mut payload = HashMap::new();
payload.insert(
"paladin_id".to_string(),
QdrantValue::from(entry.memory.paladin_id.clone()),
);
payload.insert(
"content".to_string(),
QdrantValue::from(entry.memory.content.clone()),
);
payload.insert(
"memory_type".to_string(),
QdrantValue::from(format!("{:?}", entry.memory.memory_type)),
);
payload.insert(
"importance".to_string(),
QdrantValue::from(entry.memory.importance as f64),
);
payload.insert(
"access_count".to_string(),
QdrantValue::from(entry.memory.access_count as i64),
);
payload.insert(
"created_at".to_string(),
QdrantValue::from(entry.memory.created_at.timestamp()),
);
payload.insert(
"last_accessed".to_string(),
QdrantValue::from(entry.memory.last_accessed.timestamp()),
);
for (key, value) in &entry.memory.metadata {
let qdrant_value = json_to_qdrant_value(value);
payload.insert(format!("meta_{}", key), qdrant_value);
}
Ok(PointStruct::new(
entry.memory.id.to_string(),
entry.embedding.clone(),
payload,
))
}
fn point_to_entry(
&self,
point: qdrant_client::qdrant::ScoredPoint,
) -> Result<SanctumEntry, SanctumError> {
let payload = point.payload;
let paladin_id = payload
.get("paladin_id")
.and_then(|v| v.as_str())
.ok_or_else(|| SanctumError::StorageError("Missing paladin_id".into()))?
.to_string();
let content = payload
.get("content")
.and_then(|v| v.as_str())
.ok_or_else(|| SanctumError::StorageError("Missing content".into()))?
.to_string();
let memory_type_str = payload
.get("memory_type")
.and_then(|v| v.as_str())
.ok_or_else(|| SanctumError::StorageError("Missing memory_type".into()))?;
let memory_type = match memory_type_str.as_str() {
"Episodic" => MemoryType::Episodic,
"Semantic" => MemoryType::Semantic,
"Procedural" => MemoryType::Procedural,
_ => {
return Err(SanctumError::StorageError(format!(
"Invalid memory_type: {}",
memory_type_str
)));
}
};
let importance = payload
.get("importance")
.and_then(|v| v.as_double().or_else(|| v.as_integer().map(|i| i as f64)))
.ok_or_else(|| SanctumError::StorageError("Missing importance".into()))?
as f32;
let access_count = payload
.get("access_count")
.and_then(|v| v.as_integer())
.ok_or_else(|| SanctumError::StorageError("Missing access_count".into()))?
as u32;
let created_at = payload
.get("created_at")
.and_then(|v| v.as_integer())
.ok_or_else(|| SanctumError::StorageError("Missing created_at".into()))?;
let last_accessed = payload
.get("last_accessed")
.and_then(|v| v.as_integer())
.ok_or_else(|| SanctumError::StorageError("Missing last_accessed".into()))?;
let mut metadata = HashMap::new();
for (key, value) in payload.iter() {
if let Some(meta_key) = key.strip_prefix("meta_") {
metadata.insert(meta_key.to_string(), qdrant_value_to_json(value));
}
}
let id_str = match point
.id
.as_ref()
.and_then(|id| id.point_id_options.as_ref())
{
Some(qdrant_client::qdrant::point_id::PointIdOptions::Uuid(uuid)) => uuid.clone(),
Some(qdrant_client::qdrant::point_id::PointIdOptions::Num(num)) => num.to_string(),
_ => return Err(SanctumError::StorageError("Invalid point ID".into())),
};
let id = Uuid::parse_str(&id_str)
.map_err(|e| SanctumError::StorageError(format!("Invalid UUID: {}", e)))?;
let memory = Memory {
id,
paladin_id,
content,
memory_type,
importance,
access_count,
created_at: chrono::DateTime::from_timestamp(created_at, 0)
.ok_or_else(|| SanctumError::StorageError("Invalid timestamp".into()))?,
last_accessed: chrono::DateTime::from_timestamp(last_accessed, 0)
.ok_or_else(|| SanctumError::StorageError("Invalid timestamp".into()))?,
metadata,
};
let vector = match &point.vectors {
None => return Err(SanctumError::StorageError("No vectors in point".into())),
Some(v) => match &v.vectors_options {
None => {
return Err(SanctumError::StorageError(
"No vector options in point".into(),
));
}
Some(qdrant_client::qdrant::vectors_output::VectorsOptions::Vector(vec_output)) => {
#[allow(deprecated)]
{
vec_output.data.clone()
}
}
Some(_) => {
return Err(SanctumError::StorageError(
"Unexpected vector format".into(),
));
}
},
};
SanctumEntry::new(memory, vector).map_err(SanctumError::StorageError)
}
fn build_qdrant_filter(&self, filter: &SanctumFilter) -> Option<Filter> {
let mut conditions = Vec::new();
if let Some(ref paladin_id) = filter.paladin_id {
conditions.push(Condition::matches("paladin_id", paladin_id.clone()));
}
if let Some(memory_type) = filter.memory_type {
conditions.push(Condition::matches(
"memory_type",
format!("{:?}", memory_type),
));
}
if let Some(min_importance) = filter.min_importance {
conditions.push(Condition::range(
"importance",
Range {
gte: Some(min_importance as f64),
..Default::default()
},
));
}
if let Some(created_after) = filter.created_after {
conditions.push(Condition::range(
"created_at",
Range {
gte: Some(created_after.timestamp() as f64),
..Default::default()
},
));
}
if let Some(created_before) = filter.created_before {
conditions.push(Condition::range(
"created_at",
Range {
lte: Some(created_before.timestamp() as f64),
..Default::default()
},
));
}
for (key, value) in &filter.metadata_filters {
if let Some(match_value) = json_to_match_value(value) {
conditions.push(Condition::matches(format!("meta_{}", key), match_value));
}
}
if conditions.is_empty() {
None
} else {
Some(Filter::must(conditions))
}
}
}
#[async_trait]
impl SanctumPort for QdrantSanctumAdapter {
async fn store(&self, entry: SanctumEntry) -> Result<(), SanctumError> {
let point = self.entry_to_point(&entry)?;
self.client
.upsert_points(qdrant_client::qdrant::UpsertPointsBuilder::new(
&self.collection_name,
vec![point],
))
.await
.map_err(|e| SanctumError::StorageError(format!("Failed to store entry: {}", e)))?;
Ok(())
}
async fn store_batch(&self, entries: Vec<SanctumEntry>) -> Result<(), SanctumError> {
let points: Result<Vec<_>, _> = entries.iter().map(|e| self.entry_to_point(e)).collect();
let points = points?;
self.client
.upsert_points(qdrant_client::qdrant::UpsertPointsBuilder::new(
&self.collection_name,
points,
))
.await
.map_err(|e| SanctumError::StorageError(format!("Failed to store batch: {}", e)))?;
Ok(())
}
async fn search(&self, query: SanctumQuery) -> Result<Vec<SanctumSearchResult>, SanctumError> {
if query.embedding.len() != self.vector_dimension {
return Err(SanctumError::InvalidDimension(format!(
"Expected {} dimensions, got {}",
self.vector_dimension,
query.embedding.len()
)));
}
let mut search_builder = qdrant_client::qdrant::SearchPointsBuilder::new(
&self.collection_name,
query.embedding,
query.top_k as u64,
)
.with_payload(true)
.with_vectors(true);
if let Some(min_score) = query.min_score {
search_builder = search_builder.score_threshold(min_score);
}
if let Some(ref filter) = query.filter
&& let Some(qdrant_filter) = self.build_qdrant_filter(filter)
{
search_builder = search_builder.filter(qdrant_filter);
}
let search_result = self
.client
.search_points(search_builder)
.await
.map_err(|e| SanctumError::SearchError(format!("Search failed: {}", e)))?;
let results: Result<Vec<_>, _> = search_result
.result
.into_iter()
.map(|point| {
let score = point.score;
let entry = self.point_to_entry(point)?;
Ok(SanctumSearchResult { entry, score })
})
.collect();
results
}
async fn delete(&self, id: &str) -> Result<bool, SanctumError> {
let uuid = Uuid::parse_str(id)
.map_err(|e| SanctumError::NotFound(format!("Invalid UUID: {}", e)))?;
let point_ids: Vec<_> = vec![uuid.to_string()]
.into_iter()
.map(|s| s.into())
.collect();
let result = self
.client
.get_points(
qdrant_client::qdrant::GetPointsBuilder::new(&self.collection_name, point_ids)
.with_payload(false),
)
.await
.map_err(|e| SanctumError::StorageError(format!("Failed to check existence: {}", e)))?;
if result.result.is_empty() {
return Ok(false);
}
self.client
.delete_points(
qdrant_client::qdrant::DeletePointsBuilder::new(&self.collection_name)
.points(vec![uuid.to_string()]),
)
.await
.map_err(|e| SanctumError::StorageError(format!("Failed to delete entry: {}", e)))?;
Ok(true)
}
async fn update(&self, entry: SanctumEntry) -> Result<(), SanctumError> {
self.store(entry).await
}
async fn count(&self, filter: Option<SanctumFilter>) -> Result<usize, SanctumError> {
let mut count_builder =
qdrant_client::qdrant::CountPointsBuilder::new(&self.collection_name);
if let Some(f) = filter
&& let Some(qdrant_filter) = self.build_qdrant_filter(&f)
{
count_builder = count_builder.filter(qdrant_filter);
}
let count_result = self
.client
.count(count_builder)
.await
.map_err(|e| SanctumError::StorageError(format!("Failed to count: {}", e)))?;
Ok(count_result.result.unwrap().count as usize)
}
}
fn json_to_qdrant_value(value: &Value) -> QdrantValue {
match value {
Value::String(s) => QdrantValue::from(s.clone()),
Value::Number(n) => {
if let Some(i) = n.as_i64() {
QdrantValue::from(i)
} else if let Some(f) = n.as_f64() {
QdrantValue::from(f)
} else {
QdrantValue::from(0)
}
}
Value::Bool(b) => QdrantValue::from(*b),
_ => QdrantValue::from(value.to_string()),
}
}
fn json_to_match_value(value: &Value) -> Option<MatchValue> {
match value {
Value::String(s) => Some(MatchValue::from(s.clone())),
Value::Number(n) => n.as_i64().map(MatchValue::from),
Value::Bool(b) => Some(MatchValue::from(*b)),
_ => None,
}
}
fn qdrant_value_to_json(value: &QdrantValue) -> Value {
if let Some(s) = value.as_str() {
Value::String(s.to_string())
} else if let Some(i) = value.as_integer() {
Value::Number(i.into())
} else if let Some(f) = value.as_double() {
Value::Number(serde_json::Number::from_f64(f).unwrap_or_else(|| 0.into()))
} else if let Some(b) = value.as_bool() {
Value::Bool(b)
} else {
Value::String(format!("{:?}", value))
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
use paladin_core::platform::container::sanctum::MemoryBuilder;
use paladin_ports::output::sanctum_port::{SanctumFilter, SanctumQuery};
use serde_json::json;
use std::collections::HashMap;
fn create_test_memory(paladin_id: &str, content: &str, importance: f32) -> Memory {
MemoryBuilder::new(paladin_id.to_string(), content.to_string())
.importance(importance)
.memory_type(MemoryType::Semantic)
.build()
.unwrap()
}
fn create_test_entry(paladin_id: &str, content: &str, importance: f32) -> SanctumEntry {
let memory = create_test_memory(paladin_id, content, importance);
SanctumEntry::new(memory, vec![0.1, 0.2, 0.3, 0.4, 0.5]).unwrap()
}
#[test]
fn test_sanctum_entry_creation() {
let entry = create_test_entry("paladin-1", "Test memory", 0.8);
assert_eq!(entry.paladin_id(), "paladin-1");
assert_eq!(entry.memory.content, "Test memory");
assert_eq!(entry.memory.importance, 0.8);
assert_eq!(entry.dimension, 5);
}
#[test]
fn test_memory_builder_with_metadata() {
let mut metadata = HashMap::new();
metadata.insert("source".to_string(), json!("conversation"));
let memory = MemoryBuilder::new("paladin-1".to_string(), "Test content".to_string())
.importance(0.9)
.memory_type(MemoryType::Episodic)
.metadata(metadata)
.build()
.unwrap();
assert_eq!(memory.metadata.get("source"), Some(&json!("conversation")));
}
#[test]
fn test_sanctum_entry_dimension_validation() {
let memory = create_test_memory("paladin-1", "Valid entry", 0.8);
let embedding = vec![0.1, 0.2, 0.3];
let entry = SanctumEntry::new(memory.clone(), embedding);
assert!(entry.is_ok());
let empty_embedding: Vec<f32> = vec![];
let entry = SanctumEntry::new(memory, empty_embedding);
assert!(entry.is_err());
}
#[test]
fn test_sanctum_entry_normalized_embedding() {
let memory = create_test_memory("paladin-1", "Test", 0.5);
let embedding = vec![1.0, 1.0, 1.0, 1.0];
let entry = SanctumEntry::new(memory, embedding).unwrap();
assert_eq!(entry.embedding.len(), 4);
assert_eq!(entry.dimension, 4);
}
#[test]
fn test_memory_type_variants() {
let types = vec![
MemoryType::Episodic,
MemoryType::Semantic,
MemoryType::Procedural,
];
for memory_type in types {
let memory = MemoryBuilder::new("paladin-1".to_string(), "Test".to_string())
.memory_type(memory_type)
.build()
.unwrap();
assert_eq!(memory.memory_type, memory_type);
}
}
#[test]
fn test_memory_importance_bounds() {
let valid_importances = vec![0.0, 0.5, 1.0];
for importance in valid_importances {
let memory = MemoryBuilder::new("paladin-1".to_string(), "Test".to_string())
.importance(importance)
.build();
assert!(memory.is_ok());
assert_eq!(memory.unwrap().importance, importance);
}
let memory = MemoryBuilder::new("paladin-1".to_string(), "Test".to_string())
.importance(1.5)
.build();
if let Ok(m) = memory {
assert!(m.importance <= 1.0, "Importance should be clamped to 1.0")
}
}
#[test]
fn test_sanctum_filter_builder() {
let filter = SanctumFilter::new().paladin_id("paladin-1".to_string());
assert!(filter.paladin_id.is_some());
assert_eq!(filter.paladin_id.unwrap(), "paladin-1");
let filter = SanctumFilter::new().memory_type(MemoryType::Semantic);
assert!(filter.memory_type.is_some());
assert_eq!(filter.memory_type.unwrap(), MemoryType::Semantic);
let filter = SanctumFilter::new().min_importance(0.7);
assert!(filter.min_importance.is_some());
assert_eq!(filter.min_importance.unwrap(), 0.7);
let filter = SanctumFilter::new()
.paladin_id("paladin-1".to_string())
.memory_type(MemoryType::Episodic)
.min_importance(0.8);
assert_eq!(filter.paladin_id.unwrap(), "paladin-1");
assert_eq!(filter.memory_type.unwrap(), MemoryType::Episodic);
assert_eq!(filter.min_importance.unwrap(), 0.8);
}
#[test]
fn test_sanctum_query_builder() {
let embedding = vec![0.1, 0.2, 0.3];
let query = SanctumQuery::new(embedding.clone(), 5);
assert_eq!(query.embedding.len(), 3);
assert_eq!(query.top_k, 5);
assert!(query.filter.is_none());
assert!(query.min_score.is_none());
let filter = SanctumFilter::new().paladin_id("paladin-1".to_string());
let query = SanctumQuery::new(embedding.clone(), 10).with_filter(filter);
assert!(query.filter.is_some());
assert_eq!(query.filter.unwrap().paladin_id.unwrap(), "paladin-1");
let query = SanctumQuery::new(embedding.clone(), 5).with_min_score(0.7);
assert_eq!(query.min_score.unwrap(), 0.7);
let filter = SanctumFilter::new().memory_type(MemoryType::Semantic);
let query = SanctumQuery::new(embedding, 3)
.with_filter(filter)
.with_min_score(0.85);
assert_eq!(query.top_k, 3);
assert_eq!(query.min_score.unwrap(), 0.85);
assert!(query.filter.is_some());
}
#[test]
fn test_memory_content_validation() {
let memory =
MemoryBuilder::new("paladin-1".to_string(), "Valid content".to_string()).build();
assert!(memory.is_ok());
let memory = MemoryBuilder::new("paladin-1".to_string(), "".to_string()).build();
if let Ok(m) = memory {
assert!(m.content.is_empty())
}
}
#[test]
fn test_sanctum_entry_id_uniqueness() {
let entry1 = create_test_entry("paladin-1", "First entry", 0.8);
let entry2 = create_test_entry("paladin-1", "Second entry", 0.8);
assert_ne!(
entry1.id(),
entry2.id(),
"Each entry should have a unique ID"
);
}
#[test]
fn test_memory_metadata_operations() {
let mut metadata = HashMap::new();
metadata.insert("key1".to_string(), json!("value1"));
metadata.insert("key2".to_string(), json!(42));
metadata.insert("key3".to_string(), json!(true));
let memory = MemoryBuilder::new("paladin-1".to_string(), "Test".to_string())
.metadata(metadata)
.build()
.unwrap();
assert_eq!(memory.metadata.len(), 3);
assert_eq!(memory.metadata.get("key1"), Some(&json!("value1")));
assert_eq!(memory.metadata.get("key2"), Some(&json!(42)));
assert_eq!(memory.metadata.get("key3"), Some(&json!(true)));
}
#[test]
fn test_memory_timestamps() {
let memory = create_test_memory("paladin-1", "Test", 0.5);
assert!(memory.created_at <= Utc::now());
}
}