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    /// 创建关键词包含评测器。
21    pub fn new(keywords: Vec<String>) -> Self {
22        Self {
23            keywords,
24            case_sensitive: false,
25            all_required: true,
26        }
27    }
28
29    /// 大小写敏感(默认不敏感)
30    pub fn case_sensitive(mut self, v: bool) -> Self {
31        self.case_sensitive = v;
32        self
33    }
34
35    /// true=全部包含才得分(默认);false=包含任一即得分
36    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
84/// 正则匹配评测器:检查预测是否匹配正则。
85pub struct RegexMatch {
86    pattern: Regex,
87}
88
89impl RegexMatch {
90    /// 创建正则匹配评测器(正则非法返回 `ParseError`)。
91    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
121/// 长度检查评测器:预测长度(按字符)是否落在 [min, max] 范围。
122pub 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    /// 创建长度检查评测器(默认不限长度)。
135    pub fn new() -> Self {
136        Self {
137            min: None,
138            max: None,
139        }
140    }
141
142    /// 最小长度(含)
143    pub fn min(mut self, m: usize) -> Self {
144        self.min = Some(m);
145        self
146    }
147
148    /// 最大长度(含)
149    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); // 缺"法国"
186    }
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); // 2 字符,在 [2,5]
218        assert_eq!(ev.eval("", "你好世界测试", "").await.unwrap().value, 0.0); // 6 字符,超 5
219    }
220
221    #[tokio::test]
222    async fn test_length_default_passes() {
223        // 不设 min/max,任何长度都通过
224        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}