Skip to main content

wm_memory/
query_planner.rs

1//! Deterministic query-class planner for v6 episodic retrieval.
2
3use crate::episodic_keys::{KeyCategory, extract_episodic_keys};
4
5/// Retrieval class used to select bounded scoring signals.
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum QueryClass {
8    ExactFact,
9    Temporal,
10    KnowledgeUpdate,
11    MultiHop,
12    Preference,
13    Procedure,
14    Summary,
15}
16
17/// Planner output: class plus bounded retrieval knobs.
18#[derive(Debug, Clone, PartialEq)]
19pub struct QueryPlan {
20    pub class: QueryClass,
21    pub candidate_limit: usize,
22    pub key_weight: f32,
23    pub number_query: bool,
24}
25
26impl QueryPlan {
27    /// Classify a query and return retrieval knobs.
28    #[must_use]
29    pub fn plan(query: &str, requested_limit: usize) -> Self {
30        let class = classify_query(query);
31        let requested = requested_limit.max(1);
32        let lower = query.to_ascii_lowercase();
33        let number_query = contains_any(&lower, &["how many", "how much", "how long"]);
34        match class {
35            QueryClass::ExactFact => Self {
36                class,
37                candidate_limit: requested.saturating_mul(2).max(20),
38                key_weight: 0.18,
39                number_query,
40            },
41            QueryClass::Temporal | QueryClass::KnowledgeUpdate => Self {
42                class,
43                candidate_limit: requested.saturating_mul(3).max(30),
44                key_weight: 0.15,
45                number_query,
46            },
47            QueryClass::MultiHop => Self {
48                class,
49                candidate_limit: requested.saturating_mul(4).max(40),
50                key_weight: 0.12,
51                number_query,
52            },
53            QueryClass::Preference => Self {
54                class,
55                candidate_limit: requested.saturating_mul(3).max(30),
56                key_weight: 0.25,
57                number_query,
58            },
59            QueryClass::Procedure => Self {
60                class,
61                candidate_limit: requested.saturating_mul(3).max(24),
62                key_weight: 0.1,
63                number_query,
64            },
65            QueryClass::Summary => Self {
66                class,
67                candidate_limit: requested.saturating_mul(5).max(50),
68                key_weight: 0.05,
69                number_query,
70            },
71        }
72    }
73}
74
75fn classify_query(query: &str) -> QueryClass {
76    let lower = query.to_ascii_lowercase();
77    let keys = extract_episodic_keys(query);
78    if contains_any(
79        &lower,
80        &[
81            "how do i",
82            "how to",
83            "procedure",
84            "failed",
85            "error",
86            "workaround",
87        ],
88    ) {
89        return QueryClass::Procedure;
90    }
91    if contains_any(
92        &lower,
93        &[
94            "now",
95            "currently",
96            "updated",
97            "changed",
98            "instead of",
99            "anymore",
100            "latest",
101        ],
102    ) {
103        return QueryClass::KnowledgeUpdate;
104    }
105    if keys
106        .iter()
107        .any(|key| key.category == KeyCategory::Preference)
108        || contains_any(&lower, &["favorite", "prefer", "like", "enjoy"])
109    {
110        return QueryClass::Preference;
111    }
112    if keys.iter().any(|key| key.category == KeyCategory::Date)
113        || contains_any(
114            &lower,
115            &[
116                "when",
117                "how long",
118                "before",
119                "after",
120                "last year",
121                "yesterday",
122            ],
123        )
124    {
125        return QueryClass::Temporal;
126    }
127    if contains_any(
128        &lower,
129        &[
130            "how many", "both", "and also", "across", "each of", "together",
131        ],
132    ) {
133        return QueryClass::MultiHop;
134    }
135    if contains_any(&lower, &["summarize", "overview", "all of", "everything"]) {
136        return QueryClass::Summary;
137    }
138    QueryClass::ExactFact
139}
140
141fn contains_any(haystack: &str, needles: &[&str]) -> bool {
142    needles.iter().any(|needle| haystack.contains(needle))
143}
144
145#[cfg(test)]
146mod tests {
147    use super::*;
148
149    #[test]
150    fn classifies_temporal_and_preference_queries() {
151        assert_eq!(
152            QueryPlan::plan("When did I volunteer at the animal shelter?", 10).class,
153            QueryClass::Temporal
154        );
155        assert_eq!(
156            QueryPlan::plan("What is my favorite streaming service?", 10).class,
157            QueryClass::Preference
158        );
159    }
160
161    #[test]
162    fn classifies_count_as_multihop_and_fact_as_exact() {
163        assert_eq!(
164            QueryPlan::plan("How many bikes do I own?", 10).class,
165            QueryClass::MultiHop
166        );
167        let plan = QueryPlan::plan("What degree did I graduate with?", 10);
168        assert_eq!(plan.class, QueryClass::ExactFact);
169        assert!(plan.candidate_limit >= 20);
170    }
171
172    #[test]
173    fn update_and_procedure_have_distinct_classes() {
174        assert_eq!(
175            QueryPlan::plan("What is my current internet plan now?", 10).class,
176            QueryClass::KnowledgeUpdate
177        );
178        assert_eq!(
179            QueryPlan::plan("How do I recover from a failed bike repair?", 10).class,
180            QueryClass::Procedure
181        );
182    }
183}