Skip to main content

tokio_prompt_orchestrator/
response_validator.rs

1//! # Response Validator
2//!
3//! Validates LLM responses against a set of declarative [`ValidationRule`]s.
4//! No external regex crate is required; the [`ValidationRule::MatchesRegex`]
5//! variant implements simple `*`-wildcard glob matching internally.
6//!
7//! ## Quick Start
8//!
9//! ```rust
10//! use tokio_prompt_orchestrator::response_validator::{ResponseValidator, ValidationRule};
11//!
12//! let validator = ResponseValidator::new();
13//! let rules = vec![
14//!     ValidationRule::MinLength(10),
15//!     ValidationRule::MaxLength(500),
16//!     ValidationRule::ContainsAll(vec!["Rust".to_string()]),
17//! ];
18//! let result = validator.validate("Rust is a systems programming language.", &rules);
19//! assert!(result.passed);
20//! ```
21
22use std::ops::Range;
23
24// ---------------------------------------------------------------------------
25// ValidationRule
26// ---------------------------------------------------------------------------
27
28/// A single constraint that a response must satisfy.
29#[derive(Debug, Clone)]
30pub enum ValidationRule {
31    /// Response must be no longer than `usize` characters.
32    MaxLength(usize),
33    /// Response must be at least `usize` characters long.
34    MinLength(usize),
35    /// All listed keywords must appear verbatim in the response.
36    ContainsAll(Vec<String>),
37    /// None of the listed keywords may appear in the response.
38    ContainsNone(Vec<String>),
39    /// Response must match the glob-style pattern (`*` matches any substring).
40    MatchesRegex(String),
41    /// Response must be valid JSON (parseable by `serde_json`).
42    IsValidJson,
43    /// Response must contain at least one fenced code block (``` … ```).
44    HasCodeBlock,
45    /// Number of sentences in the response must fall within the given range.
46    SentenceCount(Range<usize>),
47}
48
49impl ValidationRule {
50    /// Short human-readable name for this rule, used in violation messages.
51    pub fn name(&self) -> &str {
52        match self {
53            Self::MaxLength(_)      => "MaxLength",
54            Self::MinLength(_)      => "MinLength",
55            Self::ContainsAll(_)    => "ContainsAll",
56            Self::ContainsNone(_)   => "ContainsNone",
57            Self::MatchesRegex(_)   => "MatchesRegex",
58            Self::IsValidJson       => "IsValidJson",
59            Self::HasCodeBlock      => "HasCodeBlock",
60            Self::SentenceCount(_)  => "SentenceCount",
61        }
62    }
63}
64
65// ---------------------------------------------------------------------------
66// ValidationViolation
67// ---------------------------------------------------------------------------
68
69/// A single failed constraint.
70#[derive(Debug, Clone)]
71pub struct ValidationViolation {
72    /// The name of the rule that was violated.
73    pub rule_name: String,
74    /// Human-readable explanation of what failed and why.
75    pub message: String,
76}
77
78// ---------------------------------------------------------------------------
79// ValidationResult
80// ---------------------------------------------------------------------------
81
82/// Aggregate outcome of validating a single response.
83#[derive(Debug, Clone)]
84pub struct ValidationResult {
85    /// `true` when every rule passed (i.e. `violations` is empty).
86    pub passed: bool,
87    /// Every rule that failed, in evaluation order.
88    pub violations: Vec<ValidationViolation>,
89}
90
91impl ValidationResult {
92    fn new(violations: Vec<ValidationViolation>) -> Self {
93        Self {
94            passed: violations.is_empty(),
95            violations,
96        }
97    }
98}
99
100// ---------------------------------------------------------------------------
101// ResponseValidator
102// ---------------------------------------------------------------------------
103
104/// Validates LLM responses against a list of [`ValidationRule`]s.
105pub struct ResponseValidator;
106
107impl Default for ResponseValidator {
108    fn default() -> Self {
109        Self::new()
110    }
111}
112
113impl ResponseValidator {
114    /// Create a new validator.
115    pub fn new() -> Self {
116        Self
117    }
118
119    /// Validate `response` against every rule in `rules`.
120    ///
121    /// All rules are evaluated even when earlier ones fail, so the full list of
122    /// violations is always available.
123    pub fn validate(&self, response: &str, rules: &[ValidationRule]) -> ValidationResult {
124        let mut violations = Vec::new();
125
126        for rule in rules {
127            if let Some(v) = self.check_rule(response, rule) {
128                violations.push(v);
129            }
130        }
131
132        ValidationResult::new(violations)
133    }
134
135    /// Validate each response in `responses` against `rules`.
136    ///
137    /// Returns one [`ValidationResult`] per response in the same order.
138    pub fn validate_all(
139        &self,
140        responses: &[&str],
141        rules: &[ValidationRule],
142    ) -> Vec<ValidationResult> {
143        responses.iter().map(|r| self.validate(r, rules)).collect()
144    }
145
146    /// Fraction of results that passed all rules.
147    ///
148    /// Returns `0.0` when `results` is empty.
149    pub fn pass_rate(&self, results: &[ValidationResult]) -> f64 {
150        if results.is_empty() {
151            return 0.0;
152        }
153        let passed = results.iter().filter(|r| r.passed).count();
154        passed as f64 / results.len() as f64
155    }
156
157    // ── Per-rule check ────────────────────────────────────────────────────
158
159    fn check_rule(&self, response: &str, rule: &ValidationRule) -> Option<ValidationViolation> {
160        match rule {
161            ValidationRule::MaxLength(max) => {
162                let len = response.chars().count();
163                if len > *max {
164                    Some(ValidationViolation {
165                        rule_name: rule.name().to_string(),
166                        message: format!(
167                            "Response length {len} chars exceeds maximum {max} chars."
168                        ),
169                    })
170                } else {
171                    None
172                }
173            }
174
175            ValidationRule::MinLength(min) => {
176                let len = response.chars().count();
177                if len < *min {
178                    Some(ValidationViolation {
179                        rule_name: rule.name().to_string(),
180                        message: format!(
181                            "Response length {len} chars is below minimum {min} chars."
182                        ),
183                    })
184                } else {
185                    None
186                }
187            }
188
189            ValidationRule::ContainsAll(keywords) => {
190                let missing: Vec<&str> = keywords
191                    .iter()
192                    .filter(|kw| !response.contains(kw.as_str()))
193                    .map(|kw| kw.as_str())
194                    .collect();
195                if missing.is_empty() {
196                    None
197                } else {
198                    Some(ValidationViolation {
199                        rule_name: rule.name().to_string(),
200                        message: format!(
201                            "Response is missing required keyword(s): {}",
202                            missing.join(", ")
203                        ),
204                    })
205                }
206            }
207
208            ValidationRule::ContainsNone(keywords) => {
209                let found: Vec<&str> = keywords
210                    .iter()
211                    .filter(|kw| response.contains(kw.as_str()))
212                    .map(|kw| kw.as_str())
213                    .collect();
214                if found.is_empty() {
215                    None
216                } else {
217                    Some(ValidationViolation {
218                        rule_name: rule.name().to_string(),
219                        message: format!(
220                            "Response contains blocked keyword(s): {}",
221                            found.join(", ")
222                        ),
223                    })
224                }
225            }
226
227            ValidationRule::MatchesRegex(pattern) => {
228                if glob_match(pattern, response) {
229                    None
230                } else {
231                    Some(ValidationViolation {
232                        rule_name: rule.name().to_string(),
233                        message: format!("Response does not match glob pattern: {pattern}"),
234                    })
235                }
236            }
237
238            ValidationRule::IsValidJson => {
239                if serde_json::from_str::<serde_json::Value>(response).is_ok() {
240                    None
241                } else {
242                    Some(ValidationViolation {
243                        rule_name: rule.name().to_string(),
244                        message: "Response is not valid JSON.".to_string(),
245                    })
246                }
247            }
248
249            ValidationRule::HasCodeBlock => {
250                if response.contains("```") {
251                    None
252                } else {
253                    Some(ValidationViolation {
254                        rule_name: rule.name().to_string(),
255                        message: "Response does not contain a fenced code block (```)."
256                            .to_string(),
257                    })
258                }
259            }
260
261            ValidationRule::SentenceCount(range) => {
262                let count = count_sentences(response);
263                if range.contains(&count) {
264                    None
265                } else {
266                    Some(ValidationViolation {
267                        rule_name: rule.name().to_string(),
268                        message: format!(
269                            "Sentence count {count} is outside expected range {}..{}.",
270                            range.start, range.end
271                        ),
272                    })
273                }
274            }
275        }
276    }
277}
278
279// ---------------------------------------------------------------------------
280// Helpers
281// ---------------------------------------------------------------------------
282
283/// Simple `*`-wildcard glob match: `*` matches zero or more characters.
284///
285/// The match is case-sensitive and operates on Unicode scalar values.
286fn glob_match(pattern: &str, text: &str) -> bool {
287    // Recursive implementation over character slices.
288    let pat: Vec<char> = pattern.chars().collect();
289    let txt: Vec<char> = text.chars().collect();
290    glob_match_inner(&pat, &txt)
291}
292
293fn glob_match_inner(pat: &[char], txt: &[char]) -> bool {
294    match (pat.first(), txt.first()) {
295        // Both exhausted — matched.
296        (None, None) => true,
297        // Pattern exhausted but text remains — no match.
298        (None, Some(_)) => false,
299        // `*` in pattern: try matching zero chars (advance pattern only)
300        // or one char (advance text only).
301        (Some('*'), _) => {
302            glob_match_inner(&pat[1..], txt)
303                || (!txt.is_empty() && glob_match_inner(pat, &txt[1..]))
304        }
305        // Text exhausted but pattern still has non-`*` chars — no match.
306        (Some(_), None) => false,
307        // Literal match: both chars must be equal.
308        (Some(p), Some(t)) => {
309            if p == t {
310                glob_match_inner(&pat[1..], &txt[1..])
311            } else {
312                false
313            }
314        }
315    }
316}
317
318/// Count approximate sentence boundaries (`.`, `!`, `?`).
319///
320/// Returns at least 1 when the text is non-empty.
321fn count_sentences(text: &str) -> usize {
322    if text.trim().is_empty() {
323        return 0;
324    }
325    let count = text.chars().filter(|&c| c == '.' || c == '!' || c == '?').count();
326    count.max(1)
327}
328
329// ---------------------------------------------------------------------------
330// Tests
331// ---------------------------------------------------------------------------
332
333#[cfg(test)]
334mod tests {
335    use super::*;
336
337    fn v() -> ResponseValidator {
338        ResponseValidator::new()
339    }
340
341    // ── MaxLength ─────────────────────────────────────────────────────────
342
343    #[test]
344    fn max_length_passes_when_within() {
345        let r = v().validate("hello", &[ValidationRule::MaxLength(10)]);
346        assert!(r.passed);
347    }
348
349    #[test]
350    fn max_length_fails_when_exceeded() {
351        let r = v().validate("hello world", &[ValidationRule::MaxLength(5)]);
352        assert!(!r.passed);
353        assert_eq!(r.violations[0].rule_name, "MaxLength");
354    }
355
356    // ── MinLength ─────────────────────────────────────────────────────────
357
358    #[test]
359    fn min_length_passes_when_above() {
360        let r = v().validate("hello world", &[ValidationRule::MinLength(5)]);
361        assert!(r.passed);
362    }
363
364    #[test]
365    fn min_length_fails_when_below() {
366        let r = v().validate("hi", &[ValidationRule::MinLength(10)]);
367        assert!(!r.passed);
368        assert_eq!(r.violations[0].rule_name, "MinLength");
369    }
370
371    // ── ContainsAll ───────────────────────────────────────────────────────
372
373    #[test]
374    fn contains_all_passes_when_present() {
375        let r = v().validate(
376            "Rust is fast and safe.",
377            &[ValidationRule::ContainsAll(vec!["Rust".to_string(), "fast".to_string()])],
378        );
379        assert!(r.passed);
380    }
381
382    #[test]
383    fn contains_all_fails_when_missing() {
384        let r = v().validate(
385            "Rust is great.",
386            &[ValidationRule::ContainsAll(vec!["Python".to_string()])],
387        );
388        assert!(!r.passed);
389    }
390
391    // ── ContainsNone ──────────────────────────────────────────────────────
392
393    #[test]
394    fn contains_none_passes_when_absent() {
395        let r = v().validate(
396            "This is fine.",
397            &[ValidationRule::ContainsNone(vec!["bad".to_string()])],
398        );
399        assert!(r.passed);
400    }
401
402    #[test]
403    fn contains_none_fails_when_found() {
404        let r = v().validate(
405            "This is bad content.",
406            &[ValidationRule::ContainsNone(vec!["bad".to_string()])],
407        );
408        assert!(!r.passed);
409        assert_eq!(r.violations[0].rule_name, "ContainsNone");
410    }
411
412    // ── MatchesRegex (glob) ───────────────────────────────────────────────
413
414    #[test]
415    fn glob_star_matches_anything() {
416        let r = v().validate(
417            "anything at all",
418            &[ValidationRule::MatchesRegex("*".to_string())],
419        );
420        assert!(r.passed);
421    }
422
423    #[test]
424    fn glob_prefix_match() {
425        let r = v().validate(
426            "Hello world",
427            &[ValidationRule::MatchesRegex("Hello*".to_string())],
428        );
429        assert!(r.passed);
430    }
431
432    #[test]
433    fn glob_no_match() {
434        let r = v().validate(
435            "Goodbye world",
436            &[ValidationRule::MatchesRegex("Hello*".to_string())],
437        );
438        assert!(!r.passed);
439        assert_eq!(r.violations[0].rule_name, "MatchesRegex");
440    }
441
442    #[test]
443    fn glob_suffix_match() {
444        let r = v().validate(
445            "Hello world",
446            &[ValidationRule::MatchesRegex("*world".to_string())],
447        );
448        assert!(r.passed);
449    }
450
451    #[test]
452    fn glob_infix_match() {
453        let r = v().validate(
454            "Hello beautiful world",
455            &[ValidationRule::MatchesRegex("Hello*world".to_string())],
456        );
457        assert!(r.passed);
458    }
459
460    // ── IsValidJson ───────────────────────────────────────────────────────
461
462    #[test]
463    fn valid_json_passes() {
464        let r = v().validate(r#"{"key": "value"}"#, &[ValidationRule::IsValidJson]);
465        assert!(r.passed);
466    }
467
468    #[test]
469    fn invalid_json_fails() {
470        let r = v().validate("not json", &[ValidationRule::IsValidJson]);
471        assert!(!r.passed);
472        assert_eq!(r.violations[0].rule_name, "IsValidJson");
473    }
474
475    // ── HasCodeBlock ──────────────────────────────────────────────────────
476
477    #[test]
478    fn has_code_block_passes() {
479        let r = v().validate(
480            "Here is code:\n```rust\nfn main() {}\n```",
481            &[ValidationRule::HasCodeBlock],
482        );
483        assert!(r.passed);
484    }
485
486    #[test]
487    fn no_code_block_fails() {
488        let r = v().validate("No code here.", &[ValidationRule::HasCodeBlock]);
489        assert!(!r.passed);
490        assert_eq!(r.violations[0].rule_name, "HasCodeBlock");
491    }
492
493    // ── SentenceCount ─────────────────────────────────────────────────────
494
495    #[test]
496    fn sentence_count_in_range_passes() {
497        // "Hello. World." → 2 sentences
498        let r = v().validate(
499            "Hello. World.",
500            &[ValidationRule::SentenceCount(1..4)],
501        );
502        assert!(r.passed);
503    }
504
505    #[test]
506    fn sentence_count_out_of_range_fails() {
507        // 3 sentences, range 1..2 excludes 3
508        let r = v().validate(
509            "One. Two. Three.",
510            &[ValidationRule::SentenceCount(1..2)],
511        );
512        assert!(!r.passed);
513        assert_eq!(r.violations[0].rule_name, "SentenceCount");
514    }
515
516    // ── Multiple rules ────────────────────────────────────────────────────
517
518    #[test]
519    fn multiple_rules_all_pass() {
520        let r = v().validate(
521            "Rust is great!",
522            &[
523                ValidationRule::MinLength(5),
524                ValidationRule::MaxLength(100),
525                ValidationRule::ContainsAll(vec!["Rust".to_string()]),
526            ],
527        );
528        assert!(r.passed);
529        assert!(r.violations.is_empty());
530    }
531
532    #[test]
533    fn multiple_rules_collect_all_violations() {
534        let r = v().validate(
535            "hi",
536            &[
537                ValidationRule::MinLength(10),  // fails
538                ValidationRule::ContainsAll(vec!["Rust".to_string()]),  // fails
539            ],
540        );
541        assert!(!r.passed);
542        assert_eq!(r.violations.len(), 2);
543    }
544
545    // ── validate_all ──────────────────────────────────────────────────────
546
547    #[test]
548    fn validate_all_returns_correct_count() {
549        let validator = v();
550        let responses = ["short", "also short", "fine length response here"];
551        let rules = [ValidationRule::MinLength(10)];
552        let results = validator.validate_all(&responses, &rules);
553        assert_eq!(results.len(), 3);
554    }
555
556    // ── pass_rate ─────────────────────────────────────────────────────────
557
558    #[test]
559    fn pass_rate_all_pass() {
560        let validator = v();
561        let results = vec![
562            ValidationResult::new(vec![]),
563            ValidationResult::new(vec![]),
564        ];
565        assert!((validator.pass_rate(&results) - 1.0).abs() < f64::EPSILON);
566    }
567
568    #[test]
569    fn pass_rate_half_pass() {
570        let validator = v();
571        let results = vec![
572            ValidationResult::new(vec![]),
573            ValidationResult::new(vec![ValidationViolation {
574                rule_name: "MaxLength".to_string(),
575                message: "too long".to_string(),
576            }]),
577        ];
578        assert!((validator.pass_rate(&results) - 0.5).abs() < f64::EPSILON);
579    }
580
581    #[test]
582    fn pass_rate_empty_returns_zero() {
583        let validator = v();
584        assert_eq!(validator.pass_rate(&[]), 0.0);
585    }
586
587    // ── Default impl ──────────────────────────────────────────────────────
588
589    #[test]
590    fn default_impl_works() {
591        let v = ResponseValidator::default();
592        let r = v.validate("test", &[]);
593        assert!(r.passed);
594    }
595
596    // ── rule.name() covers all variants ──────────────────────────────────
597
598    #[test]
599    fn rule_name_non_empty_for_all_variants() {
600        let rules: Vec<ValidationRule> = vec![
601            ValidationRule::MaxLength(10),
602            ValidationRule::MinLength(1),
603            ValidationRule::ContainsAll(vec![]),
604            ValidationRule::ContainsNone(vec![]),
605            ValidationRule::MatchesRegex("*".to_string()),
606            ValidationRule::IsValidJson,
607            ValidationRule::HasCodeBlock,
608            ValidationRule::SentenceCount(0..5),
609        ];
610        for rule in &rules {
611            assert!(!rule.name().is_empty(), "{rule:?} has empty name");
612        }
613    }
614}