1use std::fmt;
8
9#[derive(Debug, Clone, PartialEq, Eq, Hash)]
13pub enum ResponseCategory {
14 Factual,
16 Opinion,
18 Creative,
20 Procedural,
22 Conversational,
24 Technical,
26 Mathematical,
28 Refusal,
30 ErrorResponse,
32 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#[derive(Debug, Clone, PartialEq, Eq, Hash)]
58pub enum QualityDimension {
59 Coherence,
61 Completeness,
63 Accuracy,
65 Conciseness,
67 Relevance,
69 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#[derive(Debug, Clone)]
91pub struct QualityScore {
92 pub dimension: QualityDimension,
94 pub score: f64,
96 pub confidence: f64,
98 pub reasoning: String,
100}
101
102#[derive(Debug, Clone)]
106pub struct ClassificationResult {
107 pub category: ResponseCategory,
109 pub confidence: f64,
111 pub quality_scores: Vec<QualityScore>,
113 pub overall_quality: f64,
115 pub is_refusal: bool,
117 pub has_citations: bool,
119 pub word_count: usize,
121 pub reading_grade: f64,
123}
124
125#[derive(Debug, Default)]
129pub struct ResponseClassifier;
130
131impl ResponseClassifier {
132 pub fn new() -> Self {
134 Self
135 }
136
137 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 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 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 pub fn detect_category(text: &str) -> (ResponseCategory, f64) {
217 let lower = text.to_lowercase();
218
219 let mut scores: Vec<(ResponseCategory, f64)> = Vec::new();
221
222 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 if Self::detect_refusal(text) {
235 return (ResponseCategory::Refusal, 0.95);
236 }
237
238 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 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 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 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 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 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 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 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 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 pub fn detect_citations(text: &str) -> bool {
328 let lower = text.to_lowercase();
329 let has_numeric_ref = text.contains("[1]") || text.contains("[2]") || text.contains("[3]");
331 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 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 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 pub fn score_completeness(response: &str, prompt: &str) -> f64 {
375 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 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 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 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 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 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; 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 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 pub fn batch_classify(&self, responses: &[(String, String)]) -> Vec<ClassificationResult> {
472 responses.iter().map(|(resp, prompt)| self.classify(resp, prompt)).collect()
473 }
474}
475
476fn 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 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
519fn 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 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}