1use crate::episodic_keys::{KeyCategory, extract_episodic_keys};
4
5#[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#[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 #[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}