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