1use std::collections::HashMap;
4use std::path::Path;
5use std::sync::Mutex;
6use std::time::SystemTime;
7
8use lru::LruCache;
9
10use crate::error::SessionStoreError;
11use crate::sessions_root;
12
13const MANIFEST_CACHE_CAPACITY: usize = 200;
14const MANIFEST_CACHE_NONZERO_CAPACITY: std::num::NonZeroUsize =
15 std::num::NonZeroUsize::new(MANIFEST_CACHE_CAPACITY).unwrap();
16static MANIFEST_CACHE: std::sync::OnceLock<Mutex<LruCache<String, CachedManifest>>> = std::sync::OnceLock::new();
17
18#[derive(Debug, Clone)]
19struct CachedManifest {
20 summary: SessionSummary,
21 signature: ManifestSignature,
22}
23
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
25struct ManifestSignature {
26 modified: Option<SystemTime>,
27 len: u64,
28}
29
30fn manifest_signature(path: &Path) -> Option<ManifestSignature> {
31 let metadata = std::fs::symlink_metadata(path).ok()?;
32 Some(ManifestSignature {
33 modified: metadata.modified().ok(),
34 len: metadata.len(),
35 })
36}
37
38pub(crate) fn invalidate_manifest_cache(path: &Path) {
40 if let Some(cache) = MANIFEST_CACHE.get()
41 && let Ok(mut cache) = cache.lock()
42 {
43 let key = path.to_string_lossy();
44 cache.pop(key.as_ref());
45 }
46}
47
48#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
50pub struct SessionSummary {
51 pub session_id: String,
53 pub turn_count: u64,
55 pub event_count: u64,
57 pub status: String,
59 pub updated_at: String,
61}
62
63#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
65pub struct FactRecord {
66 pub fact: String,
68 pub session_id: String,
70}
71
72#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq)]
78pub struct MemorySearchResult {
79 chunk_id: String,
81 path: String,
83 start_line: usize,
85 end_line: usize,
87 score: f64,
89 snippet: String,
91 source: String,
93 created_at: Option<i64>,
95}
96
97#[must_use]
99pub fn recent_sessions(workspace: &Path, n: usize) -> Vec<SessionSummary> {
100 let root = sessions_root(workspace);
101 if !root.exists() {
102 return Vec::new();
103 }
104 let mut out = Vec::new();
105 let entries = match std::fs::read_dir(&root) {
106 Ok(e) => e,
107 Err(_) => return Vec::new(),
108 };
109 let mut cache = MANIFEST_CACHE
110 .get_or_init(|| Mutex::new(LruCache::new(MANIFEST_CACHE_NONZERO_CAPACITY)))
111 .lock()
112 .unwrap_or_else(std::sync::PoisonError::into_inner);
113 for entry in entries.filter_map(Result::ok) {
114 let manifest = entry.path().join("manifest.json");
115 let key = manifest.to_string_lossy().into_owned();
116 let signature = manifest_signature(&manifest);
117 if let Some(cached) = cache.get(&key)
118 && signature == Some(cached.signature)
119 {
120 out.push(cached.summary.clone());
121 continue;
122 }
123 if let Ok(bytes) = std::fs::read(&manifest)
124 && let Ok(s) = serde_json::from_slice::<SessionSummary>(&bytes)
125 {
126 if let Some(signature) = signature.or_else(|| manifest_signature(&manifest)) {
127 cache.put(key, CachedManifest { summary: s.clone(), signature });
128 }
129 out.push(s);
130 }
131 }
132 out.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
133 out.truncate(n);
134 out
135}
136
137pub fn query_facts(workspace: &Path, limit: usize) -> Result<Vec<FactRecord>, SessionStoreError> {
141 let root = sessions_root(workspace);
142 if !root.exists() {
143 return Ok(Vec::new());
144 }
145 let mut facts: Vec<FactRecord> = Vec::new();
146 let entries = std::fs::read_dir(&root).map_err(|e| SessionStoreError::io(root.clone(), e))?;
147 for entry in entries.filter_map(Result::ok) {
148 let memory = entry.path().join(crate::DERIVED_DIR).join("memory.json");
149 let Ok(bytes) = std::fs::read(&memory) else {
150 continue;
151 };
152 let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) else {
153 continue;
154 };
155 let session_id = entry.file_name().to_string_lossy().into_owned();
156 if let Some(arr) = value.get("grounded_facts").and_then(|v| v.as_array()) {
157 for item in arr {
158 if let Some(fact) = item.get("fact").and_then(|f| f.as_str()) {
159 facts.push(FactRecord {
160 fact: fact.to_string(),
161 session_id: session_id.clone(),
162 });
163 }
164 }
165 }
166 }
167 facts.truncate(limit);
168 Ok(facts)
169}
170
171pub fn search_memory(
178 workspace: &Path,
179 query: &str,
180 max_results: usize,
181 min_score: f64,
182) -> Result<Vec<MemorySearchResult>, SessionStoreError> {
183 if query.is_empty() {
184 return Ok(Vec::new());
185 }
186
187 let root = sessions_root(workspace);
188 if !root.exists() {
189 return Ok(Vec::new());
190 }
191
192 let query_terms = tokenize(query);
193 if query_terms.is_empty() {
194 return Ok(Vec::new());
195 }
196 let mut results: Vec<MemorySearchResult> = Vec::new();
197 let session_source = String::from("session");
198 let mut documents = Vec::new();
199
200 let entries = std::fs::read_dir(&root).map_err(|e| SessionStoreError::io(root.clone(), e))?;
201 for entry in entries.filter_map(Result::ok) {
202 let session_dir = entry.path();
203 let memory = session_dir.join(crate::DERIVED_DIR).join("memory.json");
204 let Ok(bytes) = std::fs::read(&memory) else {
205 continue;
206 };
207 let Ok(value) = serde_json::from_slice::<serde_json::Value>(&bytes) else {
208 continue;
209 };
210 let session_id = entry.file_name().to_string_lossy().into_owned();
211 let memory_path = memory.to_string_lossy().into_owned();
212 let created_at = value
213 .get("created_at")
214 .and_then(|v| v.as_i64())
215 .or_else(|| value.get("updated_at").and_then(|v| v.as_i64()));
216
217 if let Some(arr) = value.get("grounded_facts").and_then(|v| v.as_array()) {
218 for (idx, item) in arr.iter().enumerate() {
219 let Some(fact) = item.get("fact").and_then(|f| f.as_str()) else {
220 continue;
221 };
222 documents.push((
223 format!("{session_id}:{idx}"),
224 memory_path.clone(),
225 fact.to_owned(),
226 tokenize(fact),
227 created_at,
228 ));
229 }
230 }
231 }
232
233 let document_count = documents.len();
234 if document_count == 0 {
235 return Ok(Vec::new());
236 }
237 let average_length =
238 documents.iter().map(|(_, _, _, terms, _)| terms.len() as f64).sum::<f64>() / document_count as f64;
239 let mut document_frequency: HashMap<String, usize> = HashMap::new();
240 for (_, _, _, terms, _) in &documents {
241 let mut seen = std::collections::HashSet::new();
242 for term in terms {
243 if seen.insert(term.as_str()) {
244 *document_frequency.entry(term.clone()).or_insert(0) += 1;
245 }
246 }
247 }
248 let k1 = 1.2;
249 let b = 0.75;
250 for (chunk_id, path, fact, terms, created_at) in documents {
251 let length = terms.len() as f64;
252 let mut term_frequency = HashMap::<&str, usize>::new();
253 for term in &terms {
254 *term_frequency.entry(term.as_str()).or_insert(0) += 1;
255 }
256 let mut score = 0.0;
257 for query_term in &query_terms {
258 let Some(&frequency) = term_frequency.get(query_term.as_str()) else {
259 continue;
260 };
261 let Some(&frequency_in_documents) = document_frequency.get(query_term) else {
262 continue;
263 };
264 let idf = (((document_count - frequency_in_documents) as f64 + 0.5)
265 / (frequency_in_documents as f64 + 0.5)
266 + 1.0)
267 .ln();
268 let denominator = frequency as f64 + k1 * (1.0 - b + b * length / average_length.max(1.0));
269 score += idf * (frequency as f64 * (k1 + 1.0)) / denominator;
270 }
271 if score <= 0.0 || score < min_score {
272 continue;
273 }
274 results.push(MemorySearchResult {
275 chunk_id,
276 path,
277 start_line: 0,
278 end_line: 0,
279 score,
280 snippet: fact,
281 source: session_source.clone(),
282 created_at,
283 });
284 }
285
286 results.sort_by(|a, b| {
287 b.score
288 .partial_cmp(&a.score)
289 .unwrap_or(std::cmp::Ordering::Equal)
290 .then_with(|| a.chunk_id.cmp(&b.chunk_id))
291 });
292 results.truncate(max_results);
293 Ok(results)
294}
295
296pub fn default_search_max_results() -> usize {
298 6
299}
300
301pub fn default_search_min_score() -> f64 {
303 0.0
304}
305
306fn count_substring_matches(text: &str, lowered_query: &str) -> usize {
307 if lowered_query.is_empty() {
308 return 0;
309 }
310 let lowered = text.to_ascii_lowercase();
311 let mut count = 0;
312 let mut start = 0;
313 while let Some(pos) = lowered[start..].find(lowered_query) {
314 count += 1;
315 start += pos + lowered_query.len();
316 }
317 count
318}
319
320fn tokenize(text: &str) -> Vec<String> {
321 let mut tokens = Vec::new();
322 let mut current = String::new();
323 for character in text.chars() {
324 if character.is_alphanumeric() {
325 current.extend(character.to_lowercase());
326 } else if !current.is_empty() {
327 tokens.push(std::mem::take(&mut current));
328 }
329 }
330 if !current.is_empty() {
331 tokens.push(current);
332 }
333 tokens
334}
335
336#[cfg(test)]
337mod tests {
338 use super::*;
339 use tempfile::TempDir;
340
341 #[test]
342 fn count_substring_matches_counts_overlapping() {
343 assert_eq!(count_substring_matches("aaaa", "aa"), 2);
344 assert_eq!(count_substring_matches("ababa", "aba"), 1);
345 assert_eq!(count_substring_matches("hello world", "ll"), 1);
346 assert_eq!(count_substring_matches("", "x"), 0);
347 }
348
349 #[test]
350 fn search_memory_returns_matching_facts() {
351 let dir = TempDir::new().expect("tempdir");
352 let sess = crate::session_dir(dir.path(), "s1");
353 std::fs::create_dir_all(sess.join(crate::DERIVED_DIR)).expect("mkdir");
354 let memory = serde_json::json!({
355 "grounded_facts": [
356 {"fact": "the widget is blue"},
357 {"fact": "the server runs on port 8080"},
358 {"fact": "use PostgreSQL for persistence"},
359 ]
360 });
361 std::fs::write(sess.join(crate::DERIVED_DIR).join("memory.json"), serde_json::to_string(&memory).expect("ser"))
362 .expect("write");
363
364 let results = search_memory(dir.path(), "blue", 10, 0.0).expect("search");
365 assert_eq!(results.len(), 1);
366 assert_eq!(results[0].snippet, "the widget is blue");
367 assert_eq!(results[0].chunk_id, "s1:0");
368 assert!(results[0].score > 0.0);
369 }
370
371 #[test]
372 fn search_memory_scores_multiple_matches() {
373 let dir = TempDir::new().expect("tempdir");
374 let sess = crate::session_dir(dir.path(), "s2");
375 std::fs::create_dir_all(sess.join(crate::DERIVED_DIR)).expect("mkdir");
376 let memory = serde_json::json!({
377 "grounded_facts": [
378 {"fact": "rust uses rustc and cargo"},
379 {"fact": "cargo is the rust build tool"},
380 ]
381 });
382 std::fs::write(sess.join(crate::DERIVED_DIR).join("memory.json"), serde_json::to_string(&memory).expect("ser"))
383 .expect("write");
384
385 let results = search_memory(dir.path(), "cargo", 10, 0.0).expect("search");
386 assert_eq!(results.len(), 2);
387 assert!(results.iter().all(|r| r.score > 0.0));
388 }
389
390 #[test]
391 fn search_memory_respects_min_score() {
392 let dir = TempDir::new().expect("tempdir");
393 let sess = crate::session_dir(dir.path(), "s3");
394 std::fs::create_dir_all(sess.join(crate::DERIVED_DIR)).expect("mkdir");
395 let memory = serde_json::json!({
396 "grounded_facts": [
397 {"fact": "alpha beta gamma"},
398 ]
399 });
400 std::fs::write(sess.join(crate::DERIVED_DIR).join("memory.json"), serde_json::to_string(&memory).expect("ser"))
401 .expect("write");
402
403 let results = search_memory(dir.path(), "beta", 10, 2.0).expect("search");
404 assert!(results.is_empty());
405 }
406
407 #[test]
408 fn search_memory_empty_query_returns_empty() {
409 let dir = TempDir::new().expect("tempdir");
410 let results = search_memory(dir.path(), "", 10, 0.0).expect("search");
411 assert!(results.is_empty());
412 }
413
414 #[test]
415 fn search_memory_sorts_by_score_descending() {
416 let dir = TempDir::new().expect("tempdir");
417 for i in 0..3 {
418 let sess = crate::session_dir(dir.path(), &format!("s{i}"));
419 std::fs::create_dir_all(sess.join(crate::DERIVED_DIR)).expect("mkdir");
420 let memory = serde_json::json!({
421 "grounded_facts": [
422 {"fact": format!("fact {i} appears twice twice")},
423 ]
424 });
425 std::fs::write(
426 sess.join(crate::DERIVED_DIR).join("memory.json"),
427 serde_json::to_string(&memory).expect("ser"),
428 )
429 .expect("write");
430 }
431
432 let results = search_memory(dir.path(), "twice", 10, 0.0).expect("search");
433 assert_eq!(results.len(), 3);
434 assert!(results.windows(2).all(|w| w[0].score >= w[1].score));
435 }
436
437 #[test]
438 fn search_memory_uses_bm25_term_coverage_and_deterministic_ties() {
439 let dir = TempDir::new().expect("tempdir");
440 for (session_id, facts) in [
441 ("s1", vec!["rust cargo tool", "unrelated note"]),
442 ("s2", vec!["cargo build tool"]),
443 ] {
444 let session = crate::session_dir(dir.path(), session_id);
445 std::fs::create_dir_all(session.join(crate::DERIVED_DIR)).expect("mkdir");
446 let memory = serde_json::json!({
447 "grounded_facts": facts.into_iter().map(|fact| serde_json::json!({"fact": fact})).collect::<Vec<_>>()
448 });
449 std::fs::write(
450 session.join(crate::DERIVED_DIR).join("memory.json"),
451 serde_json::to_string(&memory).expect("serialize"),
452 )
453 .expect("write");
454 }
455
456 let results = search_memory(dir.path(), "rust cargo", 10, 0.0).expect("search");
457 assert_eq!(results.first().map(|result| result.chunk_id.as_str()), Some("s1:0"));
458 assert!(results[0].score > results[1].score);
459 }
460
461 #[test]
462 fn recent_sessions_invalidates_manifest_cache_after_replacement() {
463 let dir = TempDir::new().expect("tempdir");
464 let session = crate::session_dir(dir.path(), "cache-session");
465 std::fs::create_dir_all(&session).expect("mkdir");
466 let manifest = serde_json::json!({
467 "session_id": "cache-session",
468 "schema_version": 1,
469 "created_at": "2026-01-01T00:00:00Z",
470 "updated_at": "2026-01-01T00:00:00Z",
471 "turn_count": 1,
472 "event_count": 1,
473 "status": "active"
474 });
475 let path = session.join("manifest.json");
476 std::fs::write(&path, serde_json::to_vec(&manifest).expect("serialize")).expect("write");
477 assert_eq!(recent_sessions(dir.path(), 1)[0].updated_at, "2026-01-01T00:00:00Z");
478
479 let mut replaced = manifest;
480 replaced["updated_at"] = serde_json::json!("2099-01-01T00:00:00Z");
481 std::fs::write(&path, serde_json::to_vec(&replaced).expect("serialize")).expect("replace");
482 assert_eq!(recent_sessions(dir.path(), 1)[0].updated_at, "2099-01-01T00:00:00Z");
483 }
484}