Skip to main content

tokio_prompt_orchestrator/
intent_classifier.rs

1//! # Intent Classifier
2//!
3//! Rule-based classifier that infers the user's intent from prompt text.
4//! Produces an [`IntentCategory`] variant and optional confidence scores.
5//!
6//! ## Quick Start
7//!
8//! ```rust
9//! use tokio_prompt_orchestrator::intent_classifier::{IntentClassifier, IntentCategory};
10//!
11//! let classifier = IntentClassifier::new();
12//! let category = classifier.classify("How do I implement a binary search in Rust?");
13//! assert_eq!(category, IntentCategory::CodeRequest);
14//! ```
15
16/// The inferred category of a user's prompt.
17#[derive(Debug, Clone, PartialEq, Eq, Hash)]
18pub enum IntentCategory {
19    /// A question seeking information or clarification.
20    Question,
21    /// A directive to perform an action.
22    Command,
23    /// A request for creative writing, storytelling, or artistic output.
24    CreativeRequest,
25    /// A request to write, explain, debug, or review code.
26    CodeRequest,
27    /// A request to analyse data, text, or a situation.
28    AnalysisRequest,
29    /// General conversational exchange.
30    Conversation,
31    /// Could not be confidently classified.
32    Unknown,
33}
34
35impl IntentCategory {
36    /// A human-readable description of this intent category.
37    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    /// The recommended model capability tier for this intent.
50    ///
51    /// Returns one of `"fast"`, `"balanced"`, or `"powerful"`.
52    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// ---------------------------------------------------------------------------
66// Feature extraction
67// ---------------------------------------------------------------------------
68
69/// Lightweight surface-level features extracted from a prompt.
70#[derive(Debug, Clone)]
71pub struct IntentFeatures {
72    /// Whether the text contains at least one `?` character.
73    pub has_question_mark: bool,
74    /// Whether the first meaningful word is an imperative verb.
75    pub starts_with_imperative: bool,
76    /// Whether any recognised code-domain keyword appears in the text.
77    pub contains_code_keywords: bool,
78    /// Whether any recognised creative-domain keyword appears in the text.
79    pub contains_creative_keywords: bool,
80    /// Approximate number of sentences (split on `.`, `!`, `?`).
81    pub sentence_count: usize,
82    /// Average word length in characters.
83    pub avg_word_length: f64,
84}
85
86// ---------------------------------------------------------------------------
87// Keyword lists
88// ---------------------------------------------------------------------------
89
90const 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
100/// Common imperative verbs that suggest a command-style intent.
101const 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
111// ---------------------------------------------------------------------------
112// Classifier
113// ---------------------------------------------------------------------------
114
115/// Rule-based intent classifier.
116///
117/// All methods take `&self` and are safe to call from multiple threads once
118/// the classifier has been constructed.
119pub struct IntentClassifier;
120
121impl Default for IntentClassifier {
122    fn default() -> Self {
123        Self::new()
124    }
125}
126
127impl IntentClassifier {
128    /// Create a new classifier with the default keyword lists.
129    pub fn new() -> Self {
130        Self
131    }
132
133    // ── Feature extraction ────────────────────────────────────────────────
134
135    /// Extract surface-level features from `text`.
136    pub fn extract_features(&self, text: &str) -> IntentFeatures {
137        let lower = text.to_lowercase();
138
139        let has_question_mark = text.contains('?');
140
141        // Sentence count: count sentence-ending punctuation characters.
142        let sentence_count = text
143            .chars()
144            .filter(|&c| c == '.' || c == '!' || c == '?')
145            .count()
146            .max(1); // treat absence of punctuation as a single sentence
147
148        // Word statistics.
149        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        // Code-keyword presence: check each word against the keyword list.
157        let contains_code_keywords = words
158            .iter()
159            .any(|w| CODE_KEYWORDS.contains(&strip_punctuation(w).as_str()));
160
161        // Also check substrings for multi-word code keywords.
162        let contains_code_keywords = contains_code_keywords
163            || CODE_KEYWORDS.iter().any(|kw| lower.contains(kw));
164
165        // Creative-keyword presence.
166        let contains_creative_keywords = CREATIVE_KEYWORDS.iter().any(|kw| lower.contains(kw));
167
168        // Imperative: does the first word match an imperative verb?
169        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    // ── Primary classify ──────────────────────────────────────────────────
188
189    /// Classify the intent of `text` into a single [`IntentCategory`].
190    ///
191    /// Uses the highest-confidence category from `classify_with_confidence`.
192    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    // ── Confidence scoring ────────────────────────────────────────────────
199
200    /// Return a score for every [`IntentCategory`], sorted highest-first.
201    ///
202    /// Scores are not probabilities; they are rule weights in `[0.0, 1.0]`.
203    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    // ── Batch classify ────────────────────────────────────────────────────
221
222    /// Classify a slice of texts, returning one [`IntentCategory`] per item.
223    pub fn batch_classify(&self, texts: &[&str]) -> Vec<IntentCategory> {
224        texts.iter().map(|t| self.classify(t)).collect()
225    }
226
227    // ── Internal scoring helpers ──────────────────────────────────────────
228
229    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        // Typical question starters (what, how, why, when, where, who, which).
235        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            // Technical prompts tend to use longer words.
267            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            // Analysis verbs are also imperatives; weigh the more specific
281            // signal above the generic imperative score (0.65) so
282            // "Analyse X" is an analysis request, not a plain command.
283            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; // small baseline
296        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
306// ---------------------------------------------------------------------------
307// Helpers
308// ---------------------------------------------------------------------------
309
310/// Remove leading/trailing punctuation from a word slice.
311fn strip_punctuation(s: &str) -> String {
312    s.trim_matches(|c: char| !c.is_alphanumeric()).to_string()
313}
314
315// ---------------------------------------------------------------------------
316// Tests
317// ---------------------------------------------------------------------------
318
319#[cfg(test)]
320mod tests {
321    use super::*;
322
323    fn classifier() -> IntentClassifier {
324        IntentClassifier::new()
325    }
326
327    // ── IntentCategory helpers ────────────────────────────────────────────
328
329    #[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    // ── Feature extraction ────────────────────────────────────────────────
366
367    #[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    // ── classify ─────────────────────────────────────────────────────────
424
425    #[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    // ── classify_with_confidence ──────────────────────────────────────────
471
472    #[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    // ── batch_classify ────────────────────────────────────────────────────
489
490    #[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    // ── Default impl ──────────────────────────────────────────────────────
518
519    #[test]
520    fn default_impl_works() {
521        let c = IntentClassifier::default();
522        let _ = c.classify("test");
523    }
524}