use std::collections::HashSet;
use async_trait::async_trait;
use chrono::Utc;
use xz_memory_core::{
Entry, EntryStore, IndexSearcher, QueryOptions, SearchOptions, SortOrder, StoreError, TimeRange,
};
use super::traits::ProjectMemory;
impl<S: EntryStore, I: IndexSearcher> super::default::LayeredMemory<S, I> {
fn project_partition(project_id: &str) -> String {
format!("project:{}", project_id)
}
}
#[async_trait]
impl<S: EntryStore, I: IndexSearcher> ProjectMemory for super::default::LayeredMemory<S, I> {
async fn put(&self, project_id: &str, key: &str, value: &str) -> Result<(), StoreError> {
let entry_id = format!("{}:{}", project_id, key);
let _ = self.store.delete(&entry_id).await;
let entry = Entry {
id: entry_id,
partition: Self::project_partition(project_id),
body: value.to_string(),
recorded_at: Utc::now().timestamp_millis() as u64,
};
self.store.append(entry).await
}
async fn get(&self, project_id: &str, key: &str) -> Result<Option<String>, StoreError> {
let partition = Self::project_partition(project_id);
let entry_id = format!("{}:{}", project_id, key);
let opts = QueryOptions { limit: usize::MAX, sort: SortOrder::Descending };
let range = TimeRange { start: None, end: None };
let entries = self.store.query(&partition, &range, &opts).await?;
let entry = entries.into_iter().find(|e| e.id == entry_id);
Ok(entry.map(|e| e.body))
}
async fn keys(&self, project_id: &str) -> Result<Vec<String>, StoreError> {
let partition = Self::project_partition(project_id);
let prefix = format!("{}:", project_id);
let opts = QueryOptions { limit: usize::MAX, sort: SortOrder::Ascending };
let range = TimeRange { start: None, end: None };
let entries = self.store.query(&partition, &range, &opts).await?;
let seen: HashSet<String> = entries
.into_iter()
.filter_map(|e| e.id.strip_prefix(&prefix).map(String::from))
.collect();
let mut keys: Vec<String> = seen.into_iter().collect();
keys.sort();
Ok(keys)
}
async fn search(
&self,
project_id: &str,
query: &str,
) -> Result<Vec<(String, f32)>, StoreError> {
let partition = Self::project_partition(project_id);
let opts = SearchOptions { limit: 10, min_relevance: None };
let results = self.searcher.search(&[partition], query, &opts).await?;
Ok(results.into_iter().map(|se| (se.entry.id, se.relevance)).collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use xz_memory_core::ScoredEntry;
use crate::backends::InMemoryEntryStore;
struct MockSearcher;
#[async_trait]
impl IndexSearcher for MockSearcher {
async fn search(
&self,
_partitions: &[String],
query: &str,
_opts: &SearchOptions,
) -> Result<Vec<ScoredEntry>, StoreError> {
Ok(vec![ScoredEntry {
entry: Entry {
id: query.to_string(),
partition: "project:p1".into(),
body: format!("match for {}", query),
recorded_at: 1000,
},
relevance: 0.9,
}])
}
}
fn setup() -> Arc<super::super::default::LayeredMemory<InMemoryEntryStore, MockSearcher>> {
Arc::new(super::super::default::LayeredMemory::new(
Arc::new(InMemoryEntryStore::new()),
Arc::new(MockSearcher),
))
}
#[tokio::test]
async fn test_put_and_get() {
let memory = setup();
memory.put("p1", "key1", "val1").await.unwrap();
memory.put("p1", "key2", "val2").await.unwrap();
assert_eq!(memory.get("p1", "key1").await.unwrap(), Some("val1".into()));
assert_eq!(memory.get("p1", "key2").await.unwrap(), Some("val2".into()));
}
#[tokio::test]
async fn test_get_nonexistent() {
let memory = setup();
assert_eq!(memory.get("p1", "missing").await.unwrap(), None);
}
#[tokio::test]
async fn test_overwrite_value() {
let memory = setup();
memory.put("p1", "key1", "v1").await.unwrap();
memory.put("p1", "key1", "v2").await.unwrap();
assert_eq!(memory.get("p1", "key1").await.unwrap(), Some("v2".into()));
}
#[tokio::test]
async fn test_keys() {
let memory = setup();
memory.put("p1", "b", "val").await.unwrap();
memory.put("p1", "a", "val").await.unwrap();
memory.put("p1", "c", "val").await.unwrap();
let keys = memory.keys("p1").await.unwrap();
assert_eq!(keys, vec!["a", "b", "c"]);
}
#[tokio::test]
async fn test_keys_empty() {
let memory = setup();
let keys = memory.keys("empty-proj").await.unwrap();
assert!(keys.is_empty());
}
#[tokio::test]
async fn test_search() {
let memory = setup();
let results = memory.search("p1", "hello").await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, "hello");
assert!((results[0].1 - 0.9).abs() < 0.01);
}
#[tokio::test]
async fn test_isolated_projects() {
let memory = setup();
memory.put("p1", "key1", "v1").await.unwrap();
memory.put("p2", "key1", "v2").await.unwrap();
assert_eq!(memory.get("p1", "key1").await.unwrap(), Some("v1".into()));
assert_eq!(memory.get("p2", "key1").await.unwrap(), Some("v2".into()));
}
}