Skip to main content

lc_evaluation/
rules.rs

1//! 规则类评测器:基于确定性规则打分(无 LLM、无向量,0 成本)。
2//!
3//! 适合答案空间明确、可机器判定的任务(格式校验、关键词、长度约束等)。
4
5use async_trait::async_trait;
6use regex::Regex;
7
8use super::{EvalError, Evaluator, Score};
9
10/// 关键词包含评测器:检查预测是否包含指定关键词。
11///
12/// `all_required = true`(默认)需全部包含才得分;`false` 包含任一即得分。
13pub 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    /// 大小写敏感(默认不敏感)
29    pub fn case_sensitive(mut self, v: bool) -> Self {
30        self.case_sensitive = v;
31        self
32    }
33
34    /// true=全部包含才得分(默认);false=包含任一即得分
35    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
83/// 正则匹配评测器:检查预测是否匹配正则。
84pub 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
119/// 长度检查评测器:预测长度(按字符)是否落在 [min, max] 范围。
120pub 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    /// 最小长度(含)
140    pub fn min(mut self, m: usize) -> Self {
141        self.min = Some(m);
142        self
143    }
144
145    /// 最大长度(含)
146    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); // 缺"法国"
183    }
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); // 2 字符,在 [2,5]
215        assert_eq!(ev.eval("", "你好世界测试", "").await.unwrap().value, 0.0); // 6 字符,超 5
216    }
217
218    #[tokio::test]
219    async fn test_length_default_passes() {
220        // 不设 min/max,任何长度都通过
221        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}