tokio_prompt_orchestrator/
response_validator.rs1use std::ops::Range;
23
24#[derive(Debug, Clone)]
30pub enum ValidationRule {
31 MaxLength(usize),
33 MinLength(usize),
35 ContainsAll(Vec<String>),
37 ContainsNone(Vec<String>),
39 MatchesRegex(String),
41 IsValidJson,
43 HasCodeBlock,
45 SentenceCount(Range<usize>),
47}
48
49impl ValidationRule {
50 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#[derive(Debug, Clone)]
71pub struct ValidationViolation {
72 pub rule_name: String,
74 pub message: String,
76}
77
78#[derive(Debug, Clone)]
84pub struct ValidationResult {
85 pub passed: bool,
87 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
100pub struct ResponseValidator;
106
107impl Default for ResponseValidator {
108 fn default() -> Self {
109 Self::new()
110 }
111}
112
113impl ResponseValidator {
114 pub fn new() -> Self {
116 Self
117 }
118
119 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 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 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 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
279fn glob_match(pattern: &str, text: &str) -> bool {
287 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 (None, None) => true,
297 (None, Some(_)) => false,
299 (Some('*'), _) => {
302 glob_match_inner(&pat[1..], txt)
303 || (!txt.is_empty() && glob_match_inner(pat, &txt[1..]))
304 }
305 (Some(_), None) => false,
307 (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
318fn 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#[cfg(test)]
334mod tests {
335 use super::*;
336
337 fn v() -> ResponseValidator {
338 ResponseValidator::new()
339 }
340
341 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[test]
496 fn sentence_count_in_range_passes() {
497 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 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 #[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), ValidationRule::ContainsAll(vec!["Rust".to_string()]), ],
540 );
541 assert!(!r.passed);
542 assert_eq!(r.violations.len(), 2);
543 }
544
545 #[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 #[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 #[test]
590 fn default_impl_works() {
591 let v = ResponseValidator::default();
592 let r = v.validate("test", &[]);
593 assert!(r.passed);
594 }
595
596 #[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}