1use async_trait::async_trait;
6use regex::Regex;
7
8use super::{EvalError, Evaluator, Score};
9
10pub struct ContainsKeyword {
14 keywords: Vec<String>,
15 case_sensitive: bool,
16 all_required: bool,
17}
18
19impl ContainsKeyword {
20 pub fn new(keywords: Vec<String>) -> Self {
21 Self {
22 keywords,
23 case_sensitive: false,
24 all_required: true,
25 }
26 }
27
28 pub fn case_sensitive(mut self, v: bool) -> Self {
30 self.case_sensitive = v;
31 self
32 }
33
34 pub fn all_required(mut self, v: bool) -> Self {
36 self.all_required = v;
37 self
38 }
39}
40
41#[async_trait]
42impl Evaluator for ContainsKeyword {
43 async fn eval(
44 &self,
45 _input: &str,
46 prediction: &str,
47 _reference: &str,
48 ) -> Result<Score, EvalError> {
49 let pred = if self.case_sensitive {
50 prediction.to_string()
51 } else {
52 prediction.to_lowercase()
53 };
54 let matches: Vec<bool> = self
55 .keywords
56 .iter()
57 .map(|k| {
58 let k = if self.case_sensitive {
59 k.clone()
60 } else {
61 k.to_lowercase()
62 };
63 pred.contains(&k)
64 })
65 .collect();
66 let ok = if self.all_required {
67 matches.iter().all(|&m| m)
68 } else {
69 matches.iter().any(|&m| m)
70 };
71 Ok(Score::new(if ok { 1.0 } else { 0.0 }).with_label(if ok {
72 "contains"
73 } else {
74 "missing"
75 }))
76 }
77
78 fn name(&self) -> &str {
79 "contains_keyword"
80 }
81}
82
83pub struct RegexMatch {
85 pattern: Regex,
86}
87
88impl RegexMatch {
89 pub fn new(pattern: &str) -> Result<Self, EvalError> {
90 Ok(Self {
91 pattern: Regex::new(pattern).map_err(|e| EvalError::ParseError(e.to_string()))?,
92 })
93 }
94}
95
96#[async_trait]
97impl Evaluator for RegexMatch {
98 async fn eval(
99 &self,
100 _input: &str,
101 prediction: &str,
102 _reference: &str,
103 ) -> Result<Score, EvalError> {
104 let ok = self.pattern.is_match(prediction);
105 Ok(
106 Score::new(if ok { 1.0 } else { 0.0 }).with_label(if ok {
107 "match"
108 } else {
109 "no_match"
110 }),
111 )
112 }
113
114 fn name(&self) -> &str {
115 "regex_match"
116 }
117}
118
119pub struct LengthCheck {
121 min: Option<usize>,
122 max: Option<usize>,
123}
124
125impl Default for LengthCheck {
126 fn default() -> Self {
127 Self::new()
128 }
129}
130
131impl LengthCheck {
132 pub fn new() -> Self {
133 Self {
134 min: None,
135 max: None,
136 }
137 }
138
139 pub fn min(mut self, m: usize) -> Self {
141 self.min = Some(m);
142 self
143 }
144
145 pub fn max(mut self, m: usize) -> Self {
147 self.max = Some(m);
148 self
149 }
150}
151
152#[async_trait]
153impl Evaluator for LengthCheck {
154 async fn eval(
155 &self,
156 _input: &str,
157 prediction: &str,
158 _reference: &str,
159 ) -> Result<Score, EvalError> {
160 let len = prediction.chars().count();
161 let ok = self.min.map_or(true, |m| len >= m) && self.max.map_or(true, |m| len <= m);
162 Ok(Score::new(if ok { 1.0 } else { 0.0 }).with_label(if ok {
163 "in_range"
164 } else {
165 "out_of_range"
166 }))
167 }
168
169 fn name(&self) -> &str {
170 "length_check"
171 }
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177
178 #[tokio::test]
179 async fn test_contains_all_required() {
180 let ev = ContainsKeyword::new(vec!["巴黎".into(), "法国".into()]);
181 assert_eq!(ev.eval("", "巴黎是法国首都", "").await.unwrap().value, 1.0);
182 assert_eq!(ev.eval("", "巴黎很大", "").await.unwrap().value, 0.0); }
184
185 #[tokio::test]
186 async fn test_contains_any() {
187 let ev = ContainsKeyword::new(vec!["巴黎".into(), "伦敦".into()]).all_required(false);
188 assert_eq!(ev.eval("", "伦敦很大", "").await.unwrap().value, 1.0);
189 assert_eq!(ev.eval("", "柏林很大", "").await.unwrap().value, 0.0);
190 }
191
192 #[tokio::test]
193 async fn test_contains_case_insensitive() {
194 let ev = ContainsKeyword::new(vec!["hello".into()]);
195 assert_eq!(ev.eval("", "Say HELLO world", "").await.unwrap().value, 1.0);
196 }
197
198 #[tokio::test]
199 async fn test_regex_match() {
200 let ev = RegexMatch::new(r"\d{4}-\d{2}-\d{2}").unwrap();
201 assert_eq!(
202 ev.eval("", "日期是 2024-01-15", "").await.unwrap().value,
203 1.0
204 );
205 assert_eq!(
206 ev.eval("", "日期是 2024/01/15", "").await.unwrap().value,
207 0.0
208 );
209 }
210
211 #[tokio::test]
212 async fn test_length_check() {
213 let ev = LengthCheck::new().min(2).max(5);
214 assert_eq!(ev.eval("", "你好", "").await.unwrap().value, 1.0); assert_eq!(ev.eval("", "你好世界测试", "").await.unwrap().value, 0.0); }
217
218 #[tokio::test]
219 async fn test_length_default_passes() {
220 let ev = LengthCheck::new();
222 assert_eq!(ev.eval("", "任意长度", "").await.unwrap().value, 1.0);
223 assert_eq!(ev.eval("", "", "").await.unwrap().value, 1.0);
224 }
225}