1#![allow(clippy::significant_drop_tightening)]
50
51use std::collections::HashMap;
52use std::sync::{Arc, Mutex};
53use std::time::Instant;
54use uuid::Uuid;
55use wm_core::{Galaxy, Result};
56
57use crate::recall::{RecallEngine, RecallResult};
58
59#[derive(Debug, Clone)]
63pub struct QueryClassification {
64 pub is_sensitive: bool,
66 pub needs_tool_calls: bool,
68 pub is_multi_turn: bool,
70 pub complexity: f32,
72 pub task_type: String,
74}
75
76impl QueryClassification {
77 #[must_use]
79 pub fn classify(query: &str) -> Self {
80 let is_sensitive = SENSITIVITY_PATTERNS
81 .iter()
82 .any(|p| query.to_lowercase().contains(p));
83
84 let needs_tool_calls = TOOL_CALL_PATTERNS
85 .iter()
86 .any(|p| query.to_lowercase().contains(p));
87
88 let is_multi_turn = MULTI_TURN_PATTERNS
89 .iter()
90 .any(|p| query.to_lowercase().contains(p));
91
92 let word_count = query.split_whitespace().count();
93 let complexity = match word_count {
94 0..=5 => 0.1,
95 6..=15 => 0.3,
96 16..=30 => 0.5,
97 31..=60 => 0.7,
98 _ => 0.9,
99 };
100
101 let task_type = if is_sensitive {
102 "sensitive_query".to_string()
103 } else if needs_tool_calls {
104 "tool_call".to_string()
105 } else if is_multi_turn {
106 "multi_turn".to_string()
107 } else if word_count < 10 {
108 "short_query".to_string()
109 } else if word_count < 30 {
110 "medium_query".to_string()
111 } else {
112 "long_query".to_string()
113 };
114
115 Self {
116 is_sensitive,
117 needs_tool_calls,
118 is_multi_turn,
119 complexity,
120 task_type,
121 }
122 }
123}
124
125static SENSITIVITY_PATTERNS: &[&str] = &[
126 "ssn",
127 "social security",
128 "passport",
129 "password",
130 "api key",
131 "secret",
132 "token",
133 "credential",
134 "credit card",
135 "bank account",
136 "diagnosis",
137 "prescription",
138 "medical record",
139 "confidential",
140 "classified",
141];
142
143static TOOL_CALL_PATTERNS: &[&str] = &[
144 "search memory",
145 "find memory",
146 "lookup",
147 "query memory",
148 "recall",
149 "use tool",
150 "call function",
151 "invoke api",
152];
153
154static MULTI_TURN_PATTERNS: &[&str] = &[
155 "then",
156 "after that",
157 "next",
158 "subsequently",
159 "finally",
160 "step 1",
161 "step 2",
162 "phase 1",
163 "phase 2",
164];
165
166#[derive(Debug, Clone)]
170pub struct ConversationalConfig {
171 pub cache_size: usize,
173 pub snippet_length: usize,
175 pub default_limit: usize,
177 pub enable_cache: bool,
179 pub exclude_private: bool,
182}
183
184impl Default for ConversationalConfig {
185 fn default() -> Self {
186 Self {
187 cache_size: 128,
188 snippet_length: 200,
189 default_limit: 10,
190 enable_cache: true,
191 exclude_private: true,
192 }
193 }
194}
195
196impl ConversationalConfig {
197 #[must_use]
199 pub fn from_env() -> Self {
200 let mut config = Self::default();
201 if let Ok(v) = std::env::var("WM_CONVERSATIONAL_CACHE_SIZE") {
202 if let Ok(n) = v.parse::<usize>() {
203 config.cache_size = n;
204 }
205 }
206 if let Ok(v) = std::env::var("WM_CONVERSATIONAL_SNIPPET_LEN") {
207 if let Ok(n) = v.parse::<usize>() {
208 config.snippet_length = n;
209 }
210 }
211 if let Ok(v) = std::env::var("WM_CONVERSATIONAL_DEFAULT_LIMIT") {
212 if let Ok(n) = v.parse::<usize>() {
213 config.default_limit = n;
214 }
215 }
216 config
217 }
218}
219
220#[derive(Debug, Clone)]
224pub struct ConversationalResult {
225 pub memory_id: Uuid,
227 pub galaxy: Galaxy,
229 pub score: f32,
231 pub snippet: String,
233 pub tags: Vec<String>,
235 pub from_cache: bool,
237 pub latency_us: u64,
239}
240
241#[derive(Debug, Clone)]
244struct CacheEntry {
245 results: Vec<RecallResult>,
246 timestamp: Instant,
247 hit_count: u32,
248}
249
250#[derive(Debug, Clone, Default)]
254pub struct SearchMetrics {
255 pub total_queries: u64,
257 pub cache_hits: u64,
259 pub cache_misses: u64,
261 pub total_latency_us: u64,
263 pub min_latency_us: u64,
265 pub max_latency_us: u64,
267}
268
269impl SearchMetrics {
270 #[must_use]
272 pub fn cache_hit_rate(&self) -> f64 {
273 let total = self.cache_hits + self.cache_misses;
274 if total == 0 {
275 0.0
276 } else {
277 self.cache_hits as f64 / total as f64
278 }
279 }
280
281 #[must_use]
283 pub fn avg_latency_us(&self) -> f64 {
284 if self.total_queries == 0 {
285 0.0
286 } else {
287 self.total_latency_us as f64 / self.total_queries as f64
288 }
289 }
290
291 #[must_use]
293 pub fn avg_latency_ms(&self) -> f64 {
294 self.avg_latency_us() / 1000.0
295 }
296
297 #[must_use]
301 pub fn meets_latency_target(&self) -> bool {
302 self.total_queries > 0 && self.avg_latency_ms() < 50.0
303 }
304}
305
306pub struct ConversationalSearch {
313 recall: Arc<RecallEngine>,
314 config: ConversationalConfig,
315 cache: Mutex<HashMap<String, CacheEntry>>,
316 cache_order: Mutex<Vec<String>>,
317 metrics: Mutex<SearchMetrics>,
318}
319
320impl ConversationalSearch {
321 #[must_use]
323 pub fn new(recall: Arc<RecallEngine>, config: ConversationalConfig) -> Self {
324 Self {
325 recall,
326 config,
327 cache: Mutex::new(HashMap::new()),
328 cache_order: Mutex::new(Vec::new()),
329 metrics: Mutex::new(SearchMetrics::default()),
330 }
331 }
332
333 #[must_use]
335 pub fn with_defaults(recall: Arc<RecallEngine>) -> Self {
336 Self::new(recall, ConversationalConfig::default())
337 }
338
339 #[must_use]
343 pub fn search(&self, query: &str, limit: Option<usize>) -> Vec<ConversationalResult> {
344 self.search_in_galaxy(query, limit, None)
345 }
346
347 #[must_use]
349 pub fn search_in_galaxy(
350 &self,
351 query: &str,
352 limit: Option<usize>,
353 galaxy: Option<Galaxy>,
354 ) -> Vec<ConversationalResult> {
355 let start = Instant::now();
356 let effective_limit = limit.unwrap_or(self.config.default_limit);
357
358 let cache_key = format!("{query}|{galaxy:?}|{effective_limit}");
360
361 let from_cache = self.config.enable_cache && {
363 let Ok(mut cache) = self.cache.lock() else {
364 return Vec::new();
365 };
366 if let Some(entry) = cache.get_mut(&cache_key) {
367 entry.hit_count += 1;
368 entry.timestamp = Instant::now();
369 true
370 } else {
371 false
372 }
373 };
374
375 let results = if from_cache {
376 let Ok(cache) = self.cache.lock() else {
377 return Vec::new();
378 };
379 cache
380 .get(&cache_key)
381 .map(|e| e.results.clone())
382 .unwrap_or_default()
383 } else {
384 let recall_results = self.recall.hybrid_search(query, effective_limit, galaxy);
386
387 if self.config.enable_cache {
389 let entry = CacheEntry {
390 results: recall_results.clone(),
391 timestamp: Instant::now(),
392 hit_count: 1,
393 };
394 {
395 let Ok(mut cache) = self.cache.lock() else {
396 return Vec::new();
397 };
398 let Ok(mut order) = self.cache_order.lock() else {
399 return Vec::new();
400 };
401
402 while order.len() >= self.config.cache_size {
404 if let Some(oldest) = order.first() {
405 cache.remove(oldest);
406 order.remove(0);
407 } else {
408 break;
409 }
410 }
411
412 cache.insert(cache_key.clone(), entry);
413 order.push(cache_key);
414 }
415 }
416
417 if self.config.enable_cache {
419 if let Ok(mut metrics) = self.metrics.lock() {
420 metrics.cache_misses += 1;
421 }
422 }
423
424 recall_results
425 };
426
427 if from_cache {
429 if let Ok(mut metrics) = self.metrics.lock() {
430 metrics.cache_hits += 1;
431 }
432 }
433
434 let latency_us = start.elapsed().as_micros() as u64;
435
436 {
438 let Ok(mut metrics) = self.metrics.lock() else {
439 return Vec::new();
440 };
441 metrics.total_queries += 1;
442 metrics.total_latency_us += latency_us;
443 if metrics.min_latency_us == 0 || latency_us < metrics.min_latency_us {
444 metrics.min_latency_us = latency_us;
445 }
446 if latency_us > metrics.max_latency_us {
447 metrics.max_latency_us = latency_us;
448 }
449 }
450
451 results
455 .into_iter()
456 .filter(|r| {
457 !self.config.exclude_private || !self.recall.is_private(r.memory_id, r.galaxy)
458 })
459 .map(|r| {
460 let snippet = if r.content.len() > self.config.snippet_length {
461 format!("{}...", &r.content[..self.config.snippet_length])
462 } else {
463 r.content.clone()
464 };
465
466 ConversationalResult {
467 memory_id: r.memory_id,
468 galaxy: r.galaxy,
469 score: r.score,
470 snippet,
471 tags: Vec::new(),
472 from_cache,
473 latency_us,
474 }
475 })
476 .collect()
477 }
478
479 #[must_use]
481 pub fn classify(&self, query: &str) -> QueryClassification {
482 QueryClassification::classify(query)
483 }
484
485 #[must_use]
487 pub fn metrics(&self) -> SearchMetrics {
488 self.metrics.lock().map(|m| m.clone()).unwrap_or_default()
489 }
490
491 pub fn clear_cache(&self) {
493 if let Ok(mut c) = self.cache.lock() {
494 c.clear();
495 }
496 if let Ok(mut c) = self.cache_order.lock() {
497 c.clear();
498 }
499 }
500
501 #[must_use]
503 pub fn cache_len(&self) -> usize {
504 self.cache.lock().map_or(0, |c| c.len())
505 }
506
507 #[must_use]
509 pub fn recall(&self) -> &RecallEngine {
510 &self.recall
511 }
512
513 #[must_use]
515 pub const fn config(&self) -> &ConversationalConfig {
516 &self.config
517 }
518
519 pub fn store(&self, galaxy: Galaxy, memory: &crate::Memory) -> Result<()> {
521 self.recall.store_with_embedding(galaxy, memory)
522 }
523}
524
525#[cfg(test)]
528mod tests {
529 use super::*;
530 use crate::embedder::{Embedder, StubEmbedder};
531 use crate::memory::Memory;
532 use crate::recall::{RecallConfig, RecallEngine};
533 use crate::search::SearchEngine;
534 use crate::store::MemoryStore;
535 use crate::vector::VectorStore;
536 use std::sync::Arc;
537 use tempfile::tempdir;
538
539 fn setup() -> (tempfile::TempDir, ConversationalSearch) {
540 let tmp = tempdir().unwrap();
541 let store = Arc::new(MemoryStore::open_default(tmp.path()).unwrap());
542 let tantivy_path = tmp.path().join("tantivy");
543 std::fs::create_dir_all(&tantivy_path).unwrap();
544 let search = Arc::new(SearchEngine::open(&tantivy_path).unwrap());
545 let vector_store = VectorStore::new();
546 let embedder: Arc<dyn Embedder> = Arc::new(StubEmbedder::new(384));
547 let recall = RecallEngine::new(
548 store,
549 search,
550 vector_store,
551 embedder,
552 RecallConfig::default(),
553 )
554 .unwrap();
555 let conv = ConversationalSearch::with_defaults(Arc::new(recall));
556 (tmp, conv)
557 }
558
559 #[test]
562 fn classify_simple_query() {
563 let c = QueryClassification::classify("hello world");
564 assert!(!c.is_sensitive);
565 assert!(!c.needs_tool_calls);
566 assert!(!c.is_multi_turn);
567 assert!(c.complexity < 0.5);
568 }
569
570 #[test]
571 fn classify_sensitive_query() {
572 let c = QueryClassification::classify("what is my password?");
573 assert!(c.is_sensitive);
574 assert_eq!(c.task_type, "sensitive_query");
575 }
576
577 #[test]
578 fn classify_tool_call_query() {
579 let c = QueryClassification::classify("search memory for rust patterns");
580 assert!(c.needs_tool_calls);
581 }
582
583 #[test]
584 fn classify_multi_turn_query() {
585 let c = QueryClassification::classify("first do X then do Y and finally do Z");
586 assert!(c.is_multi_turn);
587 }
588
589 #[test]
590 fn classify_complex_query() {
591 let c = QueryClassification::classify(
592 "analyze the complex interdisciplinary trade-offs in this nuanced scenario \
593 with multiple competing factors and conditional dependencies that require \
594 careful consideration of various architectural patterns and their implications \
595 for distributed systems design and implementation strategy",
596 );
597 assert!(c.complexity >= 0.5);
598 }
599
600 #[test]
603 fn config_defaults() {
604 let config = ConversationalConfig::default();
605 assert_eq!(config.cache_size, 128);
606 assert_eq!(config.snippet_length, 200);
607 assert_eq!(config.default_limit, 10);
608 assert!(config.enable_cache);
609 }
610
611 #[test]
614 fn metrics_empty_defaults() {
615 let metrics = SearchMetrics::default();
616 assert_eq!(metrics.cache_hit_rate(), 0.0);
617 assert_eq!(metrics.avg_latency_us(), 0.0);
618 assert!(!metrics.meets_latency_target());
619 }
620
621 #[test]
622 fn metrics_hit_rate_calculation() {
623 let metrics = SearchMetrics {
624 total_queries: 10,
625 cache_hits: 7,
626 cache_misses: 3,
627 ..Default::default()
628 };
629 assert!((metrics.cache_hit_rate() - 0.7).abs() < 0.01);
630 }
631
632 #[test]
633 fn metrics_avg_latency() {
634 let metrics = SearchMetrics {
635 total_queries: 5,
636 total_latency_us: 250_000,
637 ..Default::default()
638 };
639 assert_eq!(metrics.avg_latency_us(), 50_000.0);
640 assert_eq!(metrics.avg_latency_ms(), 50.0);
641 }
642
643 #[test]
646 fn search_empty_returns_empty() {
647 let (_tmp, conv) = setup();
648 let results = conv.search("anything", None);
649 assert!(results.is_empty());
650 }
651
652 #[test]
653 fn search_after_store_finds_results() {
654 let (_tmp, conv) = setup();
655
656 let mem = Memory::new(
657 Galaxy::Codex,
658 "Rust programming language is fast and safe".into(),
659 )
660 .with_importance(0.8)
661 .with_tags(vec!["rust".into(), "programming".into()]);
662 conv.store(Galaxy::Codex, &mem).unwrap();
663
664 let results = conv.search("rust", None);
665 assert!(!results.is_empty());
666 assert!(results[0].snippet.contains("Rust"));
667 }
668
669 #[test]
670 fn search_snippet_truncation() {
671 let (_tmp, conv) = setup();
672
673 let long_content = "A".repeat(500);
674 let mem = Memory::new(Galaxy::Codex, long_content);
675 conv.store(Galaxy::Codex, &mem).unwrap();
676
677 let results = conv.search(&"A".repeat(500), None);
678 if !results.is_empty() {
679 assert!(results[0].snippet.len() <= 203); }
681 }
682
683 #[test]
684 fn search_cache_hit_on_repeat_query() {
685 let (_tmp, conv) = setup();
686
687 let mem = Memory::new(Galaxy::Codex, "rust programming basics".into());
688 conv.store(Galaxy::Codex, &mem).unwrap();
689
690 let results1 = conv.search("rust", None);
692 assert!(!results1.is_empty());
693 assert!(!results1[0].from_cache);
694
695 let results2 = conv.search("rust", None);
697 assert!(!results2.is_empty());
698 assert!(results2[0].from_cache);
699
700 let metrics = conv.metrics();
702 assert_eq!(metrics.cache_hits, 1);
703 assert_eq!(metrics.cache_misses, 1);
704 }
705
706 #[test]
707 fn search_galaxy_filter() {
708 let (_tmp, conv) = setup();
709
710 let mem_codex = Memory::new(Galaxy::Codex, "codex memory about rust".into());
711 let mem_research = Memory::new(Galaxy::Research, "research memory about rust".into());
712
713 conv.store(Galaxy::Codex, &mem_codex).unwrap();
714 conv.store(Galaxy::Research, &mem_research).unwrap();
715
716 let results = conv.search_in_galaxy("rust", None, Some(Galaxy::Codex));
717 assert!(!results.is_empty());
718 assert!(results.iter().all(|r| r.galaxy == Galaxy::Codex));
719 }
720
721 #[test]
722 fn search_custom_limit() {
723 let (_tmp, conv) = setup();
724
725 for i in 0..5 {
726 let mem = Memory::new(Galaxy::Codex, format!("rust memory number {i}"));
727 conv.store(Galaxy::Codex, &mem).unwrap();
728 }
729
730 let results = conv.search("rust", Some(2));
731 assert!(results.len() <= 2);
732 }
733
734 #[test]
735 fn classify_without_search() {
736 let (_tmp, conv) = setup();
737 let classification = conv.classify("what is my password?");
738 assert!(classification.is_sensitive);
739 }
740
741 #[test]
742 fn cache_clear() {
743 let (_tmp, conv) = setup();
744
745 let mem = Memory::new(Galaxy::Codex, "rust memory".into());
746 conv.store(Galaxy::Codex, &mem).unwrap();
747
748 let _ = conv.search("rust", None);
750 assert_eq!(conv.cache_len(), 1);
751
752 conv.clear_cache();
754 assert_eq!(conv.cache_len(), 0);
755 }
756
757 #[test]
758 fn cache_lru_eviction() {
759 let (_tmp, conv) = setup();
760
761 let _recall = conv.recall();
763 let small_config = ConversationalConfig {
764 cache_size: 3,
765 ..Default::default()
766 };
767 let conv_small = ConversationalSearch::new(
768 Arc::new(
769 RecallEngine::new(
770 Arc::new(
771 MemoryStore::open_default(tempfile::tempdir().unwrap().path()).unwrap(),
772 ),
773 Arc::new(SearchEngine::open(tempfile::tempdir().unwrap().path()).unwrap()),
774 VectorStore::new(),
775 Arc::new(StubEmbedder::new(384)),
776 RecallConfig::default(),
777 )
778 .unwrap(),
779 ),
780 small_config,
781 );
782
783 for i in 0..5 {
785 let _ = conv_small.search(&format!("query{i}"), None);
786 }
787
788 assert!(conv_small.cache_len() <= 3);
790 }
791
792 #[test]
793 fn metrics_track_latency() {
794 let (_tmp, conv) = setup();
795
796 let mem = Memory::new(Galaxy::Codex, "rust memory".into());
797 conv.store(Galaxy::Codex, &mem).unwrap();
798
799 let _ = conv.search("rust", None);
800
801 let metrics = conv.metrics();
802 assert!(metrics.total_queries > 0);
803 assert!(metrics.total_latency_us > 0);
804 assert!(metrics.max_latency_us >= metrics.min_latency_us);
805 }
806
807 #[test]
808 fn store_delegates_to_recall() {
809 let (_tmp, conv) = setup();
810
811 let mem = Memory::new(Galaxy::Codex, "test content for delegation".into());
812 conv.store(Galaxy::Codex, &mem).unwrap();
813
814 let results = conv.search("test content for delegation", None);
816 assert!(!results.is_empty());
817 }
818
819 #[test]
820 fn search_with_disabled_cache() {
821 let tmp = tempdir().unwrap();
822 let store = Arc::new(MemoryStore::open_default(tmp.path()).unwrap());
823 let tantivy_path = tmp.path().join("tantivy");
824 std::fs::create_dir_all(&tantivy_path).unwrap();
825 let search = Arc::new(SearchEngine::open(&tantivy_path).unwrap());
826 let vector_store = VectorStore::new();
827 let embedder: Arc<dyn Embedder> = Arc::new(StubEmbedder::new(384));
828 let recall = RecallEngine::new(
829 store,
830 search,
831 vector_store,
832 embedder,
833 RecallConfig::default(),
834 )
835 .unwrap();
836
837 let config = ConversationalConfig {
838 enable_cache: false,
839 ..Default::default()
840 };
841 let conv = ConversationalSearch::new(Arc::new(recall), config);
842
843 let mem = Memory::new(Galaxy::Codex, "rust memory".into());
844 conv.store(Galaxy::Codex, &mem).unwrap();
845
846 let _ = conv.search("rust", None);
847 let _ = conv.search("rust", None);
848
849 assert_eq!(conv.cache_len(), 0);
851
852 let metrics = conv.metrics();
853 assert_eq!(metrics.cache_hits, 0);
854 assert_eq!(metrics.cache_misses, 0);
855 }
856}