Skip to main content

tokio_prompt_orchestrator/
response_classifier.rs

1//! LLM response classification and quality scoring.
2//!
3//! Classifies free-text LLM responses into semantic categories, computes
4//! multi-dimensional quality scores, detects refusals, citations, and
5//! computes a Flesch-Kincaid readability grade level.
6
7use std::fmt;
8
9// ── ResponseCategory ──────────────────────────────────────────────────────────
10
11/// High-level semantic category of an LLM response.
12#[derive(Debug, Clone, PartialEq, Eq, Hash)]
13pub enum ResponseCategory {
14    /// The response presents objective facts.
15    Factual,
16    /// The response expresses a subjective opinion.
17    Opinion,
18    /// The response is creative writing (story, poem, etc.).
19    Creative,
20    /// The response describes a sequence of steps or a procedure.
21    Procedural,
22    /// The response is casual conversational text.
23    Conversational,
24    /// The response covers a technical or engineering topic.
25    Technical,
26    /// The response involves numerical computation or math.
27    Mathematical,
28    /// The model declined to answer.
29    Refusal,
30    /// The response is itself an error message.
31    ErrorResponse,
32    /// The category could not be determined with confidence.
33    Uncertain,
34}
35
36impl fmt::Display for ResponseCategory {
37    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
38        let s = match self {
39            ResponseCategory::Factual => "Factual",
40            ResponseCategory::Opinion => "Opinion",
41            ResponseCategory::Creative => "Creative",
42            ResponseCategory::Procedural => "Procedural",
43            ResponseCategory::Conversational => "Conversational",
44            ResponseCategory::Technical => "Technical",
45            ResponseCategory::Mathematical => "Mathematical",
46            ResponseCategory::Refusal => "Refusal",
47            ResponseCategory::ErrorResponse => "ErrorResponse",
48            ResponseCategory::Uncertain => "Uncertain",
49        };
50        write!(f, "{s}")
51    }
52}
53
54// ── QualityDimension ──────────────────────────────────────────────────────────
55
56/// An axis along which response quality is evaluated.
57#[derive(Debug, Clone, PartialEq, Eq, Hash)]
58pub enum QualityDimension {
59    /// Logical flow and internal consistency.
60    Coherence,
61    /// Whether the response fully addresses the prompt.
62    Completeness,
63    /// Factual correctness (heuristic approximation).
64    Accuracy,
65    /// Brevity vs. verbosity balance.
66    Conciseness,
67    /// How well the response stays on topic.
68    Relevance,
69    /// Use of structure (headers, lists, code blocks).
70    Formatting,
71}
72
73impl fmt::Display for QualityDimension {
74    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
75        let s = match self {
76            QualityDimension::Coherence => "Coherence",
77            QualityDimension::Completeness => "Completeness",
78            QualityDimension::Accuracy => "Accuracy",
79            QualityDimension::Conciseness => "Conciseness",
80            QualityDimension::Relevance => "Relevance",
81            QualityDimension::Formatting => "Formatting",
82        };
83        write!(f, "{s}")
84    }
85}
86
87// ── QualityScore ──────────────────────────────────────────────────────────────
88
89/// A scored measurement on a single quality dimension.
90#[derive(Debug, Clone)]
91pub struct QualityScore {
92    /// Which dimension was measured.
93    pub dimension: QualityDimension,
94    /// Score in [0.0, 1.0].
95    pub score: f64,
96    /// Confidence in the score in [0.0, 1.0].
97    pub confidence: f64,
98    /// Human-readable explanation of the score.
99    pub reasoning: String,
100}
101
102// ── ClassificationResult ──────────────────────────────────────────────────────
103
104/// The full result of classifying one LLM response.
105#[derive(Debug, Clone)]
106pub struct ClassificationResult {
107    /// Predicted semantic category.
108    pub category: ResponseCategory,
109    /// Confidence in the category prediction in [0.0, 1.0].
110    pub confidence: f64,
111    /// Per-dimension quality scores.
112    pub quality_scores: Vec<QualityScore>,
113    /// Weighted mean of all dimension scores.
114    pub overall_quality: f64,
115    /// Whether the model refused to answer.
116    pub is_refusal: bool,
117    /// Whether citations or references were detected.
118    pub has_citations: bool,
119    /// Number of words in the response.
120    pub word_count: usize,
121    /// Flesch-Kincaid grade level.
122    pub reading_grade: f64,
123}
124
125// ── ResponseClassifier ────────────────────────────────────────────────────────
126
127/// Stateless classifier for LLM responses.
128#[derive(Debug, Default)]
129pub struct ResponseClassifier;
130
131impl ResponseClassifier {
132    /// Create a new classifier.
133    pub fn new() -> Self {
134        Self
135    }
136
137    /// Classify a single response given the original prompt.
138    pub fn classify(&self, response: &str, prompt: &str) -> ClassificationResult {
139        let is_refusal = Self::detect_refusal(response);
140        let has_citations = Self::detect_citations(response);
141        let word_count = count_words(response);
142        let reading_grade = Self::flesch_kincaid_grade(response);
143
144        let (category, confidence) = if is_refusal {
145            (ResponseCategory::Refusal, 0.95)
146        } else {
147            Self::detect_category(response)
148        };
149
150        let coherence = QualityScore {
151            dimension: QualityDimension::Coherence,
152            score: Self::score_coherence(response),
153            confidence: 0.6,
154            reasoning: "Sentence transition and topic consistency heuristic.".to_string(),
155        };
156        let completeness = QualityScore {
157            dimension: QualityDimension::Completeness,
158            score: Self::score_completeness(response, prompt),
159            confidence: 0.55,
160            reasoning: "Question-word coverage relative to prompt.".to_string(),
161        };
162        let conciseness = QualityScore {
163            dimension: QualityDimension::Conciseness,
164            score: Self::score_conciseness(response),
165            confidence: 0.5,
166            reasoning: "Filler phrase and repetition penalty.".to_string(),
167        };
168        let formatting = QualityScore {
169            dimension: QualityDimension::Formatting,
170            score: Self::score_formatting(response),
171            confidence: 0.7,
172            reasoning: "Presence of lists, headers, and code blocks.".to_string(),
173        };
174        // Accuracy: proxy — higher for factual/technical, lower for uncertain.
175        let accuracy_score = match &category {
176            ResponseCategory::Factual | ResponseCategory::Technical => 0.75,
177            ResponseCategory::Mathematical => 0.80,
178            ResponseCategory::Opinion | ResponseCategory::Creative => 0.60,
179            ResponseCategory::Refusal | ResponseCategory::ErrorResponse => 0.30,
180            _ => 0.55,
181        };
182        let accuracy = QualityScore {
183            dimension: QualityDimension::Accuracy,
184            score: accuracy_score,
185            confidence: 0.4,
186            reasoning: "Category-based accuracy proxy.".to_string(),
187        };
188        // Relevance: overlap between prompt tokens and response tokens.
189        let relevance_score = compute_token_overlap(prompt, response);
190        let relevance = QualityScore {
191            dimension: QualityDimension::Relevance,
192            score: relevance_score,
193            confidence: 0.6,
194            reasoning: "Token overlap between prompt and response.".to_string(),
195        };
196
197        let quality_scores = vec![coherence, completeness, accuracy, conciseness, relevance, formatting];
198        let overall_quality = quality_scores.iter().map(|s| s.score).sum::<f64>()
199            / quality_scores.len() as f64;
200
201        ClassificationResult {
202            category,
203            confidence,
204            quality_scores,
205            overall_quality,
206            is_refusal,
207            has_citations,
208            word_count,
209            reading_grade,
210        }
211    }
212
213    /// Detect the response category using keyword/pattern matching.
214    ///
215    /// Returns `(category, confidence)`.
216    pub fn detect_category(text: &str) -> (ResponseCategory, f64) {
217        let lower = text.to_lowercase();
218
219        // Scores accumulated per category.
220        let mut scores: Vec<(ResponseCategory, f64)> = Vec::new();
221
222        // Mathematical: digits, operators, equations.
223        let math_score = {
224            let eq_count = lower.matches('=').count();
225            let digit_density = lower.chars().filter(|c| c.is_ascii_digit()).count() as f64
226                / lower.len().max(1) as f64;
227            let kw = count_keywords(&lower, &["equation", "calculate", "formula", "integral",
228                "derivative", "matrix", "theorem", "proof", "sum", "product"]);
229            (eq_count as f64 * 0.05 + digit_density * 2.0 + kw as f64 * 0.1).min(1.0)
230        };
231        scores.push((ResponseCategory::Mathematical, math_score));
232
233        // Refusal: handled before calling this function, but guard anyway.
234        if Self::detect_refusal(text) {
235            return (ResponseCategory::Refusal, 0.95);
236        }
237
238        // Error response.
239        let err_score = {
240            let kw = count_keywords(&lower, &["error:", "exception:", "traceback", "stack trace",
241                "syntax error", "runtime error", "null pointer", "segmentation fault"]);
242            (kw as f64 * 0.25).min(1.0)
243        };
244        scores.push((ResponseCategory::ErrorResponse, err_score));
245
246        // Technical.
247        let tech_score = {
248            let kw = count_keywords(&lower, &["function", "struct", "impl", "class", "module",
249                "algorithm", "api", "database", "server", "protocol", "async", "thread",
250                "memory", "cpu", "network", "interface", "library", "framework", "compile",
251                "runtime", "binary", "architecture"]);
252            let code_blocks = lower.matches("```").count() as f64;
253            (kw as f64 * 0.06 + code_blocks * 0.15).min(1.0)
254        };
255        scores.push((ResponseCategory::Technical, tech_score));
256
257        // Procedural: step-by-step, numbered lists.
258        let proc_score = {
259            let kw = count_keywords(&lower, &["step", "first", "second", "third", "next",
260                "then", "finally", "install", "configure", "run", "execute", "follow"]);
261            let numbered = lower.lines().filter(|l| {
262                let t = l.trim();
263                t.starts_with("1.") || t.starts_with("2.") || t.starts_with("3.")
264            }).count();
265            (kw as f64 * 0.05 + numbered as f64 * 0.1).min(1.0)
266        };
267        scores.push((ResponseCategory::Procedural, proc_score));
268
269        // Factual.
270        let fact_score = {
271            let kw = count_keywords(&lower, &["according to", "research shows", "studies indicate",
272                "published", "evidence", "data", "statistics", "was born", "founded in",
273                "located in", "discovered", "invented", "historically"]);
274            (kw as f64 * 0.12).min(1.0)
275        };
276        scores.push((ResponseCategory::Factual, fact_score));
277
278        // Opinion.
279        let opinion_score = {
280            let kw = count_keywords(&lower, &["i think", "i believe", "in my opinion",
281                "i feel", "personally", "i would say", "i recommend", "arguably", "seems to me"]);
282            (kw as f64 * 0.15).min(1.0)
283        };
284        scores.push((ResponseCategory::Opinion, opinion_score));
285
286        // Creative.
287        let creative_score = {
288            let kw = count_keywords(&lower, &["once upon a time", "she said", "he said",
289                "chapter", "verse", "rhyme", "stanza", "protagonist", "narrative",
290                "story", "poem", "fiction", "character"]);
291            (kw as f64 * 0.1).min(1.0)
292        };
293        scores.push((ResponseCategory::Creative, creative_score));
294
295        // Conversational: short, informal.
296        let conv_score = {
297            let word_count = count_words(text);
298            let kw = count_keywords(&lower, &["sure!", "of course", "happy to", "great question",
299                "thanks", "you're welcome", "absolutely", "definitely"]);
300            let short_bonus = if word_count < 60 { 0.2 } else { 0.0 };
301            (kw as f64 * 0.1 + short_bonus).min(1.0)
302        };
303        scores.push((ResponseCategory::Conversational, conv_score));
304
305        // Pick highest score.
306        let best = scores.into_iter().max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
307        match best {
308            Some((cat, score)) if score >= 0.15 => (cat, (score + 0.4).min(1.0)),
309            _ => (ResponseCategory::Uncertain, 0.3),
310        }
311    }
312
313    /// Detect if the response is a refusal.
314    pub fn detect_refusal(text: &str) -> bool {
315        let lower = text.to_lowercase();
316        let phrases = [
317            "i cannot", "i can't", "i am unable", "i'm unable",
318            "i won't", "i will not", "as an ai", "as a language model",
319            "i don't have the ability", "i'm not able", "i am not able",
320            "i must decline", "i refuse", "that's not something i",
321            "i cannot assist", "i'm afraid i cannot",
322        ];
323        phrases.iter().any(|p| lower.contains(p))
324    }
325
326    /// Detect whether the response contains citations or references.
327    pub fn detect_citations(text: &str) -> bool {
328        let lower = text.to_lowercase();
329        // [1] style footnotes
330        let has_numeric_ref = text.contains("[1]") || text.contains("[2]") || text.contains("[3]");
331        // "according to", "source:", "cf.", "see:"
332        let has_phrases = lower.contains("according to")
333            || lower.contains("source:")
334            || lower.contains(" cf.")
335            || lower.contains("see:")
336            || lower.contains("cited in")
337            || lower.contains("references:");
338        has_numeric_ref || has_phrases
339    }
340
341    /// Score sentence-level coherence (0-1).
342    ///
343    /// Heuristic: rewards smooth sentence beginnings, penalises abrupt stops.
344    pub fn score_coherence(text: &str) -> f64 {
345        let sentences: Vec<&str> = split_sentences(text);
346        if sentences.len() < 2 {
347            return 0.7;
348        }
349
350        let transition_words = ["however", "therefore", "furthermore", "additionally",
351            "moreover", "consequently", "thus", "hence", "also", "similarly",
352            "in contrast", "on the other hand", "for example", "as a result"];
353
354        let transitions = sentences.iter().filter(|s| {
355            let lower = s.to_lowercase();
356            transition_words.iter().any(|t| lower.starts_with(t) || lower.contains(&format!(", {t}")))
357        }).count();
358
359        let transition_rate = transitions as f64 / sentences.len() as f64;
360
361        // Average sentence length variance: very short or very long sentences hurt coherence.
362        let lengths: Vec<usize> = sentences.iter().map(|s| count_words(s)).collect();
363        let mean_len = lengths.iter().sum::<usize>() as f64 / lengths.len() as f64;
364        let variance = lengths.iter().map(|&l| {
365            let diff = l as f64 - mean_len;
366            diff * diff
367        }).sum::<f64>() / lengths.len() as f64;
368        let length_penalty = (variance / 100.0).min(0.3);
369
370        (0.5 + transition_rate * 0.5 - length_penalty).clamp(0.0, 1.0)
371    }
372
373    /// Score how completely the response addresses the prompt (0-1).
374    pub fn score_completeness(response: &str, prompt: &str) -> f64 {
375        // Check if question words in the prompt are addressed in the response.
376        let prompt_lower = prompt.to_lowercase();
377        let response_lower = response.to_lowercase();
378
379        let question_words = ["what", "why", "how", "when", "where", "who", "which", "explain"];
380        let asked: Vec<&str> = question_words.iter().filter(|w| prompt_lower.contains(*w)).copied().collect();
381
382        if asked.is_empty() {
383            // No explicit question words — check basic length adequacy.
384            let words = count_words(response);
385            return if words >= 50 { 0.8 } else { 0.5 };
386        }
387
388        let answered = asked.iter().filter(|w| response_lower.contains(*w)).count();
389        let base = answered as f64 / asked.len() as f64;
390
391        // Length bonus: very short answers likely incomplete.
392        let words = count_words(response);
393        let length_bonus = if words >= 100 { 0.15 } else if words >= 40 { 0.05 } else { 0.0 };
394
395        (base * 0.85 + length_bonus).clamp(0.0, 1.0)
396    }
397
398    /// Score conciseness (0-1). Higher = more concise.
399    pub fn score_conciseness(text: &str) -> f64 {
400        let lower = text.to_lowercase();
401        let filler_phrases = [
402            "it is important to note that",
403            "it should be noted that",
404            "in order to",
405            "due to the fact that",
406            "at this point in time",
407            "for the purpose of",
408            "in the event that",
409            "the fact that",
410            "it is worth mentioning",
411            "needless to say",
412            "as a matter of fact",
413        ];
414
415        let filler_count = filler_phrases.iter().filter(|p| lower.contains(*p)).count();
416        let filler_penalty = (filler_count as f64 * 0.08).min(0.4);
417
418        // Repetition: count how many 4-gram sequences repeat.
419        let words: Vec<&str> = text.split_whitespace().collect();
420        let mut ngrams: std::collections::HashMap<[&str; 4], usize> = std::collections::HashMap::new();
421        for w in words.windows(4) {
422            *ngrams.entry([w[0], w[1], w[2], w[3]]).or_insert(0) += 1;
423        }
424        let repeated = ngrams.values().filter(|&&c| c > 1).count();
425        let repetition_penalty = (repeated as f64 * 0.05).min(0.3);
426
427        (1.0 - filler_penalty - repetition_penalty).clamp(0.0, 1.0)
428    }
429
430    /// Score formatting quality (0-1).
431    pub fn score_formatting(text: &str) -> f64 {
432        let has_code_block = text.contains("```");
433        let has_numbered_list = text.lines().any(|l| {
434            let t = l.trim();
435            t.len() > 2 && t.chars().next().map(|c| c.is_ascii_digit()).unwrap_or(false) && t.contains(". ")
436        });
437        let has_bullet_list = text.lines().any(|l| {
438            let t = l.trim();
439            t.starts_with("- ") || t.starts_with("* ") || t.starts_with("• ")
440        });
441        let has_header = text.lines().any(|l| l.trim().starts_with('#'));
442
443        let mut score: f64 = 0.4; // baseline
444        if has_code_block { score += 0.2; }
445        if has_numbered_list { score += 0.15; }
446        if has_bullet_list { score += 0.15; }
447        if has_header { score += 0.1; }
448
449        score.clamp(0.0, 1.0)
450    }
451
452    /// Compute the Flesch-Kincaid Grade Level.
453    ///
454    /// FK Grade = 0.39 * (words/sentences) + 11.8 * (syllables/words) - 15.59
455    pub fn flesch_kincaid_grade(text: &str) -> f64 {
456        let word_count = count_words(text);
457        if word_count == 0 {
458            return 0.0;
459        }
460        let sentence_count = split_sentences(text).len().max(1);
461        let syllable_count = text.split_whitespace().map(count_syllables).sum::<usize>();
462
463        let words_per_sentence = word_count as f64 / sentence_count as f64;
464        let syllables_per_word = syllable_count as f64 / word_count as f64;
465
466        let grade = 0.39 * words_per_sentence + 11.8 * syllables_per_word - 15.59;
467        grade.clamp(0.0, 20.0)
468    }
469
470    /// Classify a batch of (response, prompt) pairs.
471    pub fn batch_classify(&self, responses: &[(String, String)]) -> Vec<ClassificationResult> {
472        responses.iter().map(|(resp, prompt)| self.classify(resp, prompt)).collect()
473    }
474}
475
476// ── Helpers ───────────────────────────────────────────────────────────────────
477
478fn count_words(text: &str) -> usize {
479    text.split_whitespace().count()
480}
481
482fn count_keywords(text: &str, keywords: &[&str]) -> usize {
483    keywords.iter().filter(|k| text.contains(*k)).count()
484}
485
486fn split_sentences(text: &str) -> Vec<&str> {
487    // Simple heuristic: split on `. `, `! `, `? `
488    let mut result = Vec::new();
489    let mut start = 0;
490    let bytes = text.as_bytes();
491    let len = bytes.len();
492    let mut i = 0;
493    while i < len {
494        if (bytes[i] == b'.' || bytes[i] == b'!' || bytes[i] == b'?')
495            && i + 1 < len && bytes[i + 1] == b' '
496        {
497            let s = text[start..=i].trim();
498            if !s.is_empty() {
499                result.push(s);
500            }
501            start = i + 2;
502            i += 2;
503        } else {
504            i += 1;
505        }
506    }
507    if start < len {
508        let s = text[start..].trim();
509        if !s.is_empty() {
510            result.push(s);
511        }
512    }
513    if result.is_empty() {
514        result.push(text.trim());
515    }
516    result
517}
518
519/// Approximate syllable count for a single word.
520fn count_syllables(word: &str) -> usize {
521    let lower = word.to_lowercase();
522    let vowels = "aeiouy";
523    let chars: Vec<char> = lower.chars().collect();
524    let mut count = 0usize;
525    let mut prev_vowel = false;
526    for &c in &chars {
527        let is_vowel = vowels.contains(c);
528        if is_vowel && !prev_vowel {
529            count += 1;
530        }
531        prev_vowel = is_vowel;
532    }
533    // Silent 'e' at end.
534    if lower.ends_with('e') && count > 1 {
535        count -= 1;
536    }
537    count.max(1)
538}
539
540fn compute_token_overlap(prompt: &str, response: &str) -> f64 {
541    let prompt_tokens: std::collections::HashSet<&str> = prompt.split_whitespace().collect();
542    let response_tokens: std::collections::HashSet<&str> = response.split_whitespace().collect();
543    if prompt_tokens.is_empty() {
544        return 0.5;
545    }
546    let overlap = prompt_tokens.intersection(&response_tokens).count();
547    (overlap as f64 / prompt_tokens.len() as f64).clamp(0.0, 1.0)
548}
549
550#[cfg(test)]
551mod tests {
552    use super::*;
553
554    #[test]
555    fn test_refusal_detection() {
556        assert!(ResponseClassifier::detect_refusal("I cannot help with that request."));
557        assert!(ResponseClassifier::detect_refusal("As an AI, I won't provide that."));
558        assert!(!ResponseClassifier::detect_refusal("The answer is 42."));
559    }
560
561    #[test]
562    fn test_citation_detection() {
563        assert!(ResponseClassifier::detect_citations("See [1] for more details."));
564        assert!(ResponseClassifier::detect_citations("According to research, this is true."));
565        assert!(!ResponseClassifier::detect_citations("The quick brown fox."));
566    }
567
568    #[test]
569    fn test_flesch_kincaid() {
570        let text = "The cat sat on the mat. It was a good cat.";
571        let grade = ResponseClassifier::flesch_kincaid_grade(text);
572        assert!(grade >= 0.0 && grade <= 20.0);
573    }
574
575    #[test]
576    fn test_classify_basic() {
577        let classifier = ResponseClassifier::new();
578        let result = classifier.classify(
579            "The function takes two arguments and returns a struct.",
580            "What does this function do?",
581        );
582        assert!(result.overall_quality > 0.0 && result.overall_quality <= 1.0);
583        assert!(result.word_count > 0);
584    }
585
586    #[test]
587    fn test_batch_classify() {
588        let classifier = ResponseClassifier::new();
589        let pairs = vec![
590            ("Hello!".to_string(), "Hi".to_string()),
591            ("The Earth is 4.5 billion years old.".to_string(), "How old is Earth?".to_string()),
592        ];
593        let results = classifier.batch_classify(&pairs);
594        assert_eq!(results.len(), 2);
595    }
596
597    #[test]
598    fn test_score_formatting_with_code() {
599        let text = "Here is the code:\n```rust\nfn main() {}\n```";
600        let score = ResponseClassifier::score_formatting(text);
601        assert!(score > 0.4);
602    }
603}