1#[derive(Debug, Clone, PartialEq, Eq, Hash)]
18pub enum IntentCategory {
19 Question,
21 Command,
23 CreativeRequest,
25 CodeRequest,
27 AnalysisRequest,
29 Conversation,
31 Unknown,
33}
34
35impl IntentCategory {
36 pub fn description(&self) -> &str {
38 match self {
39 Self::Question => "User is asking a question seeking information or clarification.",
40 Self::Command => "User is issuing a directive to perform a specific action.",
41 Self::CreativeRequest => "User wants creative writing, storytelling, or artistic content.",
42 Self::CodeRequest => "User wants code written, explained, debugged, or reviewed.",
43 Self::AnalysisRequest => "User wants data, text, or a situation analysed.",
44 Self::Conversation => "General conversational exchange with no specific task.",
45 Self::Unknown => "Intent could not be confidently determined.",
46 }
47 }
48
49 pub fn suggested_model_tier(&self) -> &str {
53 match self {
54 Self::Conversation => "fast",
55 Self::Question => "fast",
56 Self::Command => "balanced",
57 Self::CreativeRequest => "balanced",
58 Self::AnalysisRequest => "powerful",
59 Self::CodeRequest => "powerful",
60 Self::Unknown => "balanced",
61 }
62 }
63}
64
65#[derive(Debug, Clone)]
71pub struct IntentFeatures {
72 pub has_question_mark: bool,
74 pub starts_with_imperative: bool,
76 pub contains_code_keywords: bool,
78 pub contains_creative_keywords: bool,
80 pub sentence_count: usize,
82 pub avg_word_length: f64,
84}
85
86const CODE_KEYWORDS: &[&str] = &[
91 "fn", "function", "class", "def", "code", "implement", "debug",
92 "error", "compile", "rust", "python", "javascript", "algorithm",
93];
94
95const CREATIVE_KEYWORDS: &[&str] = &[
96 "write", "story", "poem", "creative", "imagine", "generate",
97 "design", "create", "art",
98];
99
100const IMPERATIVE_VERBS: &[&str] = &[
102 "list", "show", "find", "get", "set", "run", "execute", "make",
103 "build", "open", "close", "delete", "remove", "add", "start",
104 "stop", "install", "update", "check", "print", "display", "fetch",
105 "send", "move", "copy", "rename", "convert", "parse", "sort",
106 "filter", "search", "count", "calculate", "compute", "generate",
107 "create", "write", "define", "summarise", "summarize", "translate",
108 "explain", "describe", "analyse", "analyze", "compare", "classify",
109];
110
111pub struct IntentClassifier;
120
121impl Default for IntentClassifier {
122 fn default() -> Self {
123 Self::new()
124 }
125}
126
127impl IntentClassifier {
128 pub fn new() -> Self {
130 Self
131 }
132
133 pub fn extract_features(&self, text: &str) -> IntentFeatures {
137 let lower = text.to_lowercase();
138
139 let has_question_mark = text.contains('?');
140
141 let sentence_count = text
143 .chars()
144 .filter(|&c| c == '.' || c == '!' || c == '?')
145 .count()
146 .max(1); let words: Vec<&str> = lower.split_whitespace().collect();
150 let avg_word_length = if words.is_empty() {
151 0.0
152 } else {
153 words.iter().map(|w| w.len()).sum::<usize>() as f64 / words.len() as f64
154 };
155
156 let contains_code_keywords = words
158 .iter()
159 .any(|w| CODE_KEYWORDS.contains(&strip_punctuation(w).as_str()));
160
161 let contains_code_keywords = contains_code_keywords
163 || CODE_KEYWORDS.iter().any(|kw| lower.contains(kw));
164
165 let contains_creative_keywords = CREATIVE_KEYWORDS.iter().any(|kw| lower.contains(kw));
167
168 let starts_with_imperative = words
170 .first()
171 .map(|w| {
172 let bare = strip_punctuation(w);
173 IMPERATIVE_VERBS.contains(&bare.as_str())
174 })
175 .unwrap_or(false);
176
177 IntentFeatures {
178 has_question_mark,
179 starts_with_imperative,
180 contains_code_keywords,
181 contains_creative_keywords,
182 sentence_count,
183 avg_word_length,
184 }
185 }
186
187 pub fn classify(&self, text: &str) -> IntentCategory {
193 let mut ranked = self.classify_with_confidence(text);
194 ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
195 ranked.into_iter().next().map(|(cat, _)| cat).unwrap_or(IntentCategory::Unknown)
196 }
197
198 pub fn classify_with_confidence(&self, text: &str) -> Vec<(IntentCategory, f64)> {
204 let f = self.extract_features(text);
205
206 let mut scores: Vec<(IntentCategory, f64)> = vec![
207 (IntentCategory::Question, self.score_question(&f)),
208 (IntentCategory::Command, self.score_command(&f)),
209 (IntentCategory::CreativeRequest, self.score_creative(&f)),
210 (IntentCategory::CodeRequest, self.score_code(&f)),
211 (IntentCategory::AnalysisRequest, self.score_analysis(text, &f)),
212 (IntentCategory::Conversation, self.score_conversation(&f)),
213 (IntentCategory::Unknown, 0.05),
214 ];
215
216 scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
217 scores
218 }
219
220 pub fn batch_classify(&self, texts: &[&str]) -> Vec<IntentCategory> {
224 texts.iter().map(|t| self.classify(t)).collect()
225 }
226
227 fn score_question(&self, f: &IntentFeatures) -> f64 {
230 let mut score = 0.0_f64;
231 if f.has_question_mark {
232 score += 0.60;
233 }
234 score.min(1.0)
236 }
237
238 fn score_command(&self, f: &IntentFeatures) -> f64 {
239 let mut score = 0.0_f64;
240 if f.starts_with_imperative {
241 score += 0.55;
242 }
243 if !f.has_question_mark {
244 score += 0.10;
245 }
246 score.min(1.0)
247 }
248
249 fn score_creative(&self, f: &IntentFeatures) -> f64 {
250 let mut score = 0.0_f64;
251 if f.contains_creative_keywords {
252 score += 0.65;
253 }
254 if !f.contains_code_keywords {
255 score += 0.10;
256 }
257 score.min(1.0)
258 }
259
260 fn score_code(&self, f: &IntentFeatures) -> f64 {
261 let mut score = 0.0_f64;
262 if f.contains_code_keywords {
263 score += 0.70;
264 }
265 if f.avg_word_length > 5.5 {
266 score += 0.10;
268 }
269 score.min(1.0)
270 }
271
272 fn score_analysis(&self, text: &str, f: &IntentFeatures) -> f64 {
273 let lower = text.to_lowercase();
274 let mut score = 0.0_f64;
275 let analysis_words = ["analyse", "analyze", "analysis", "compare",
276 "evaluate", "assess", "review", "examine",
277 "investigate", "breakdown", "break down",
278 "summarise", "summarize", "interpret"];
279 for kw in &analysis_words {
280 if lower.contains(kw) {
284 score += 0.70;
285 break;
286 }
287 }
288 if f.sentence_count > 2 {
289 score += 0.10;
290 }
291 score.min(1.0)
292 }
293
294 fn score_conversation(&self, f: &IntentFeatures) -> f64 {
295 let mut score = 0.15_f64; if f.sentence_count == 1 && f.avg_word_length < 5.0 {
297 score += 0.30;
298 }
299 if !f.contains_code_keywords && !f.contains_creative_keywords {
300 score += 0.10;
301 }
302 score.min(1.0)
303 }
304}
305
306fn strip_punctuation(s: &str) -> String {
312 s.trim_matches(|c: char| !c.is_alphanumeric()).to_string()
313}
314
315#[cfg(test)]
320mod tests {
321 use super::*;
322
323 fn classifier() -> IntentClassifier {
324 IntentClassifier::new()
325 }
326
327 #[test]
330 fn description_non_empty_for_all_variants() {
331 let variants = [
332 IntentCategory::Question,
333 IntentCategory::Command,
334 IntentCategory::CreativeRequest,
335 IntentCategory::CodeRequest,
336 IntentCategory::AnalysisRequest,
337 IntentCategory::Conversation,
338 IntentCategory::Unknown,
339 ];
340 for v in &variants {
341 assert!(!v.description().is_empty(), "{v:?} has empty description");
342 }
343 }
344
345 #[test]
346 fn suggested_model_tier_valid_values() {
347 let valid = ["fast", "balanced", "powerful"];
348 let variants = [
349 IntentCategory::Question,
350 IntentCategory::Command,
351 IntentCategory::CreativeRequest,
352 IntentCategory::CodeRequest,
353 IntentCategory::AnalysisRequest,
354 IntentCategory::Conversation,
355 IntentCategory::Unknown,
356 ];
357 for v in &variants {
358 assert!(
359 valid.contains(&v.suggested_model_tier()),
360 "{v:?} returned invalid tier"
361 );
362 }
363 }
364
365 #[test]
368 fn extract_features_question_mark() {
369 let c = classifier();
370 let f = c.extract_features("What is Rust?");
371 assert!(f.has_question_mark);
372 }
373
374 #[test]
375 fn extract_features_no_question_mark() {
376 let c = classifier();
377 let f = c.extract_features("Tell me about Rust.");
378 assert!(!f.has_question_mark);
379 }
380
381 #[test]
382 fn extract_features_imperative() {
383 let c = classifier();
384 let f = c.extract_features("List all files in the directory.");
385 assert!(f.starts_with_imperative);
386 }
387
388 #[test]
389 fn extract_features_not_imperative() {
390 let c = classifier();
391 let f = c.extract_features("The quick brown fox.");
392 assert!(!f.starts_with_imperative);
393 }
394
395 #[test]
396 fn extract_features_code_keywords() {
397 let c = classifier();
398 let f = c.extract_features("Help me debug this rust function.");
399 assert!(f.contains_code_keywords);
400 }
401
402 #[test]
403 fn extract_features_creative_keywords() {
404 let c = classifier();
405 let f = c.extract_features("Write a poem about autumn.");
406 assert!(f.contains_creative_keywords);
407 }
408
409 #[test]
410 fn extract_features_sentence_count() {
411 let c = classifier();
412 let f = c.extract_features("First sentence. Second sentence! Third?");
413 assert_eq!(f.sentence_count, 3);
414 }
415
416 #[test]
417 fn extract_features_avg_word_length_positive() {
418 let c = classifier();
419 let f = c.extract_features("hello world");
420 assert!(f.avg_word_length > 0.0);
421 }
422
423 #[test]
426 fn classify_code_request() {
427 let c = classifier();
428 assert_eq!(
429 c.classify("How do I implement a binary search algorithm in Rust?"),
430 IntentCategory::CodeRequest
431 );
432 }
433
434 #[test]
435 fn classify_creative_request() {
436 let c = classifier();
437 assert_eq!(
438 c.classify("Write me a short story about a lonely robot."),
439 IntentCategory::CreativeRequest
440 );
441 }
442
443 #[test]
444 fn classify_question() {
445 let c = classifier();
446 assert_eq!(
447 c.classify("What is the capital of France?"),
448 IntentCategory::Question
449 );
450 }
451
452 #[test]
453 fn classify_command() {
454 let c = classifier();
455 assert_eq!(
456 c.classify("List all the environment variables."),
457 IntentCategory::Command
458 );
459 }
460
461 #[test]
462 fn classify_analysis() {
463 let c = classifier();
464 assert_eq!(
465 c.classify("Analyse the trade-offs between SQL and NoSQL databases."),
466 IntentCategory::AnalysisRequest
467 );
468 }
469
470 #[test]
473 fn confidence_sorted_descending() {
474 let c = classifier();
475 let scores = c.classify_with_confidence("How do I write a function in Python?");
476 for w in scores.windows(2) {
477 assert!(w[0].1 >= w[1].1, "scores not sorted: {:?}", scores);
478 }
479 }
480
481 #[test]
482 fn confidence_all_categories_present() {
483 let c = classifier();
484 let scores = c.classify_with_confidence("hi there");
485 assert_eq!(scores.len(), 7);
486 }
487
488 #[test]
491 fn batch_classify_length_matches() {
492 let c = classifier();
493 let texts = ["hello", "write a poem", "debug this code"];
494 let results = c.batch_classify(&texts);
495 assert_eq!(results.len(), 3);
496 }
497
498 #[test]
499 fn batch_classify_empty_input() {
500 let c = classifier();
501 let results = c.batch_classify(&[]);
502 assert!(results.is_empty());
503 }
504
505 #[test]
506 fn batch_classify_mixed_intents() {
507 let c = classifier();
508 let texts = [
509 "Implement a Rust function.",
510 "Write me a poem about the ocean.",
511 ];
512 let results = c.batch_classify(&texts);
513 assert_eq!(results[0], IntentCategory::CodeRequest);
514 assert_eq!(results[1], IntentCategory::CreativeRequest);
515 }
516
517 #[test]
520 fn default_impl_works() {
521 let c = IntentClassifier::default();
522 let _ = c.classify("test");
523 }
524}