Skip to main content

xz_memory_engine/layered/
project.rs

1//! [`ProjectMemory`] implementation for [`LayeredMemory`].
2
3use std::collections::HashSet;
4
5use async_trait::async_trait;
6use chrono::Utc;
7use xz_memory_core::{
8    Entry, EntryStore, IndexSearcher, QueryOptions, SearchOptions, SortOrder, StoreError, TimeRange,
9};
10
11use super::traits::ProjectMemory;
12
13impl<S: EntryStore, I: IndexSearcher> super::default::LayeredMemory<S, I> {
14    fn project_partition(project_id: &str) -> String {
15        format!("project:{}", project_id)
16    }
17}
18
19#[async_trait]
20impl<S: EntryStore, I: IndexSearcher> ProjectMemory for super::default::LayeredMemory<S, I> {
21    async fn put(&self, project_id: &str, key: &str, value: &str) -> Result<(), StoreError> {
22        let entry_id = format!("{}:{}", project_id, key);
23        let _ = self.store.delete(&entry_id).await;
24        let entry = Entry {
25            id: entry_id,
26            partition: Self::project_partition(project_id),
27            body: value.to_string(),
28            recorded_at: Utc::now().timestamp_millis() as u64,
29        };
30        self.store.append(entry).await
31    }
32
33    async fn get(&self, project_id: &str, key: &str) -> Result<Option<String>, StoreError> {
34        let partition = Self::project_partition(project_id);
35        let entry_id = format!("{}:{}", project_id, key);
36        let opts = QueryOptions { limit: usize::MAX, sort: SortOrder::Descending };
37        let range = TimeRange { start: None, end: None };
38        let entries = self.store.query(&partition, &range, &opts).await?;
39        let entry = entries.into_iter().find(|e| e.id == entry_id);
40        Ok(entry.map(|e| e.body))
41    }
42
43    async fn keys(&self, project_id: &str) -> Result<Vec<String>, StoreError> {
44        let partition = Self::project_partition(project_id);
45        let prefix = format!("{}:", project_id);
46        let opts = QueryOptions { limit: usize::MAX, sort: SortOrder::Ascending };
47        let range = TimeRange { start: None, end: None };
48        let entries = self.store.query(&partition, &range, &opts).await?;
49        let seen: HashSet<String> = entries
50            .into_iter()
51            .filter_map(|e| e.id.strip_prefix(&prefix).map(String::from))
52            .collect();
53        let mut keys: Vec<String> = seen.into_iter().collect();
54        keys.sort();
55        Ok(keys)
56    }
57
58    async fn search(
59        &self,
60        project_id: &str,
61        query: &str,
62    ) -> Result<Vec<(String, f32)>, StoreError> {
63        let partition = Self::project_partition(project_id);
64        let opts = SearchOptions { limit: 10, min_relevance: None };
65        let results = self.searcher.search(&[partition], query, &opts).await?;
66        Ok(results.into_iter().map(|se| (se.entry.id, se.relevance)).collect())
67    }
68}
69
70#[cfg(test)]
71mod tests {
72    use super::*;
73    use std::sync::Arc;
74    use xz_memory_core::ScoredEntry;
75
76    use crate::backends::InMemoryEntryStore;
77
78    struct MockSearcher;
79
80    #[async_trait]
81    impl IndexSearcher for MockSearcher {
82        async fn search(
83            &self,
84            _partitions: &[String],
85            query: &str,
86            _opts: &SearchOptions,
87        ) -> Result<Vec<ScoredEntry>, StoreError> {
88            Ok(vec![ScoredEntry {
89                entry: Entry {
90                    id: query.to_string(),
91                    partition: "project:p1".into(),
92                    body: format!("match for {}", query),
93                    recorded_at: 1000,
94                },
95                relevance: 0.9,
96            }])
97        }
98    }
99
100    fn setup() -> Arc<super::super::default::LayeredMemory<InMemoryEntryStore, MockSearcher>> {
101        Arc::new(super::super::default::LayeredMemory::new(
102            Arc::new(InMemoryEntryStore::new()),
103            Arc::new(MockSearcher),
104        ))
105    }
106
107    #[tokio::test]
108    async fn test_put_and_get() {
109        let memory = setup();
110
111        memory.put("p1", "key1", "val1").await.unwrap();
112        memory.put("p1", "key2", "val2").await.unwrap();
113
114        assert_eq!(memory.get("p1", "key1").await.unwrap(), Some("val1".into()));
115        assert_eq!(memory.get("p1", "key2").await.unwrap(), Some("val2".into()));
116    }
117
118    #[tokio::test]
119    async fn test_get_nonexistent() {
120        let memory = setup();
121        assert_eq!(memory.get("p1", "missing").await.unwrap(), None);
122    }
123
124    #[tokio::test]
125    async fn test_overwrite_value() {
126        let memory = setup();
127
128        memory.put("p1", "key1", "v1").await.unwrap();
129        memory.put("p1", "key1", "v2").await.unwrap();
130
131        assert_eq!(memory.get("p1", "key1").await.unwrap(), Some("v2".into()));
132    }
133
134    #[tokio::test]
135    async fn test_keys() {
136        let memory = setup();
137
138        memory.put("p1", "b", "val").await.unwrap();
139        memory.put("p1", "a", "val").await.unwrap();
140        memory.put("p1", "c", "val").await.unwrap();
141
142        let keys = memory.keys("p1").await.unwrap();
143        assert_eq!(keys, vec!["a", "b", "c"]);
144    }
145
146    #[tokio::test]
147    async fn test_keys_empty() {
148        let memory = setup();
149        let keys = memory.keys("empty-proj").await.unwrap();
150        assert!(keys.is_empty());
151    }
152
153    #[tokio::test]
154    async fn test_search() {
155        let memory = setup();
156
157        let results = memory.search("p1", "hello").await.unwrap();
158        assert_eq!(results.len(), 1);
159        assert_eq!(results[0].0, "hello");
160        assert!((results[0].1 - 0.9).abs() < 0.01);
161    }
162
163    #[tokio::test]
164    async fn test_isolated_projects() {
165        let memory = setup();
166
167        memory.put("p1", "key1", "v1").await.unwrap();
168        memory.put("p2", "key1", "v2").await.unwrap();
169
170        assert_eq!(memory.get("p1", "key1").await.unwrap(), Some("v1".into()));
171        assert_eq!(memory.get("p2", "key1").await.unwrap(), Some("v2".into()));
172    }
173}