Skip to main content

lc_evaluation/
results.rs

1//! Built-in evaluators: ExactMatch, StringDistance, EmbeddingSimilarity, LLMAsJudge.
2
3use async_trait::async_trait;
4use serde::Deserialize;
5
6use lc_core::tools::ToolDefinition;
7use lc_core::BaseChatModel;
8use lc_embeddings::{cosine_similarity, Embeddings};
9use lc_schema::Message;
10
11use lc_core::judge::{structured_call, truncate, StructuredJudgeError};
12
13use super::criteria::{EvalError, Evaluator, Score};
14
15pub struct ExactMatch;
16
17#[async_trait]
18impl Evaluator for ExactMatch {
19    async fn eval(
20        &self,
21        _input: &str,
22        prediction: &str,
23        reference: &str,
24    ) -> Result<Score, EvalError> {
25        let matched = prediction.trim() == reference.trim();
26        let v = if matched { 1.0 } else { 0.0 };
27        let label = if matched { "match" } else { "mismatch" };
28        Ok(Score::new(v).with_label(label))
29    }
30    fn name(&self) -> &str {
31        "exact_match"
32    }
33}
34
35pub struct StringDistance;
36
37impl StringDistance {
38    fn levenshtein(a: &str, b: &str) -> usize {
39        let a: Vec<char> = a.chars().collect();
40        let b: Vec<char> = b.chars().collect();
41        let (m, n) = (a.len(), b.len());
42        if m == 0 {
43            return n;
44        }
45        if n == 0 {
46            return m;
47        }
48        let mut prev: Vec<usize> = (0..=n).collect();
49        let mut curr: Vec<usize> = vec![0; n + 1];
50        for i in 1..=m {
51            curr[0] = i;
52            for j in 1..=n {
53                let cost = if a[i - 1] == b[j - 1] { 0 } else { 1 };
54                curr[j] = (prev[j] + 1).min(curr[j - 1] + 1).min(prev[j - 1] + cost);
55            }
56            std::mem::swap(&mut prev, &mut curr);
57        }
58        prev[n]
59    }
60}
61
62#[async_trait]
63impl Evaluator for StringDistance {
64    async fn eval(
65        &self,
66        _input: &str,
67        prediction: &str,
68        reference: &str,
69    ) -> Result<Score, EvalError> {
70        let dist = Self::levenshtein(prediction, reference) as f64;
71        let max_len = prediction.chars().count().max(reference.chars().count()) as f64;
72        let score = if max_len == 0.0 {
73            1.0
74        } else {
75            1.0 - dist / max_len
76        };
77        Ok(Score::new(score))
78    }
79    fn name(&self) -> &str {
80        "string_distance"
81    }
82}
83
84pub struct EmbeddingSimilarity<E: Embeddings> {
85    embeddings: E,
86}
87
88impl<E: Embeddings> EmbeddingSimilarity<E> {
89    pub fn new(embeddings: E) -> Self {
90        Self { embeddings }
91    }
92}
93
94#[async_trait]
95impl<E: Embeddings> Evaluator for EmbeddingSimilarity<E> {
96    async fn eval(
97        &self,
98        _input: &str,
99        prediction: &str,
100        reference: &str,
101    ) -> Result<Score, EvalError> {
102        let p = self
103            .embeddings
104            .embed_query(prediction)
105            .await
106            .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
107        let r = self
108            .embeddings
109            .embed_query(reference)
110            .await
111            .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
112        // P2-7: cosine 长度不匹配是数据缺陷(向量维度不一致),不再吞成 0.0 静默降级
113        let sim =
114            cosine_similarity(&p, &r).map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
115        let v = ((sim + 1.0) / 2.0).clamp(0.0, 1.0);
116        Ok(Score::new(v as f64))
117    }
118    fn name(&self) -> &str {
119        "embedding_similarity"
120    }
121}
122
123pub struct LLMAsJudge<M: BaseChatModel> {
124    judge: M,
125    rubric: String,
126    max_score: u8,
127}
128
129const DEFAULT_RUBRIC: &str = "\
130正确性:回答是否事实准确、是否与参考答案的核心意思一致。
131完整性:是否完整回答了输入的问题或指令。
132清晰性:表达是否清晰、无歧义、无冗余。";
133
134impl<M: BaseChatModel> LLMAsJudge<M> {
135    pub fn new(judge: M) -> Self {
136        Self {
137            judge,
138            rubric: DEFAULT_RUBRIC.to_string(),
139            max_score: 10,
140        }
141    }
142    pub fn with_rubric(mut self, rubric: impl Into<String>) -> Self {
143        self.rubric = rubric.into();
144        self
145    }
146    pub fn with_max_score(mut self, max_score: u8) -> Self {
147        self.max_score = max_score.max(1);
148        self
149    }
150
151    fn build_prompt(&self, input: &str, prediction: &str, reference: &str) -> (String, String) {
152        let system = format!(
153            "你是一个严格、公正的评估员。请根据以下评分标准对待评估的回答打分。\n\n评分标准:\n{rubric}\n\n打分范围:0 到 {max}(0 = 完全错误或无关,{max} = 完全正确)。\n\n要求:先在 reason 字段写出简短分析,再在 score 字段给出分数。\n只输出一行 JSON,格式为:{{\"reason\":\"...\",\"score\":N}}",
154            rubric = self.rubric, max = self.max_score
155        );
156        let user =
157            format!(
158            "输入:\n{input}\n\n参考答案:\n{reference}\n\n待评估的回答:\n{prediction}\n\n请评估。",
159            input = input, reference = reference, prediction = prediction,
160        );
161        (system, user)
162    }
163}
164
165/// 结构化评分参数(经 tool_calls 返回)。
166#[derive(Debug, Deserialize)]
167struct ScoreArgs {
168    score: f64,
169    /// 让 LLM 附上简短理由(改善打分质量),当前不消费。
170    #[serde(default)]
171    #[allow(dead_code)]
172    reason: String,
173}
174
175/// 构建评分工具:让 LLM 以 `{"score": 0..max, "reason": "..."}` 提交评分。
176fn score_tool(max_score: u8) -> ToolDefinition {
177    ToolDefinition::new(
178        "submit_evaluation",
179        "提交你对回答的评分。score 为 0 到 max 的整数,reason 给出简短分析。",
180    )
181    .with_parameters(serde_json::json!({
182        "type": "object",
183        "properties": {
184            "score": {
185                "type": "integer",
186                "minimum": 0,
187                "maximum": max_score,
188                "description": "0 到 max 的整数分数"
189            },
190            "reason": { "type": "string", "description": "简短分析" }
191        },
192        "required": ["score", "reason"]
193    }))
194}
195
196#[async_trait]
197impl<M: BaseChatModel> Evaluator for LLMAsJudge<M> {
198    async fn eval(
199        &self,
200        input: &str,
201        prediction: &str,
202        reference: &str,
203    ) -> Result<Score, EvalError> {
204        let (system, user) = self.build_prompt(input, prediction, reference);
205        let messages = vec![Message::system(system), Message::human(user)];
206
207        // P0-1: 优先结构化输出(tool_calls);不支持工具绑定的模型走文本解析回落。
208        let args: ScoreArgs =
209            structured_call(&self.judge, score_tool(self.max_score), messages, |raw| {
210                let norm = parse_score(raw, self.max_score).ok_or_else(|| {
211                    StructuredJudgeError::Parse(format!(
212                        "无法从裁判回复解析分数: {}",
213                        truncate(raw, 200)
214                    ))
215                })?;
216                Ok(ScoreArgs {
217                    score: norm * self.max_score as f64,
218                    reason: String::new(),
219                })
220            })
221            .await?;
222
223        let value = (args.score / self.max_score as f64).clamp(0.0, 1.0);
224        Ok(Score::new(value).with_label("llm_judge"))
225    }
226    fn name(&self) -> &str {
227        "llm_as_judge"
228    }
229}
230
231fn parse_score(raw: &str, max_score: u8) -> Option<f64> {
232    let max = max_score as f64;
233    let n = extract_json_score(raw)
234        .or_else(|| find_number_after_keyword(raw, "score"))
235        .or_else(|| find_number_after_keyword(raw, "分数"))
236        .or_else(|| first_number(raw))?;
237    Some((n / max).clamp(0.0, 1.0))
238}
239
240fn extract_json_score(raw: &str) -> Option<f64> {
241    let start = raw.find('{')?;
242    let end = raw.rfind('}')?;
243    if end < start {
244        return None;
245    }
246    let val: serde_json::Value = serde_json::from_str(&raw[start..=end]).ok()?;
247    val.get("score")?.as_f64()
248}
249
250fn find_number_after_keyword(raw: &str, keyword: &str) -> Option<f64> {
251    let lower_raw = raw.to_lowercase();
252    let lower_kw = keyword.to_lowercase();
253    let idx = lower_raw.find(lower_kw.as_str())?;
254    first_number(&raw[idx + lower_kw.len()..])
255}
256
257fn first_number(s: &str) -> Option<f64> {
258    let mut buf = String::new();
259    let mut started = false;
260    for c in s.chars() {
261        if c.is_ascii_digit() || c == '.' {
262            started = true;
263            buf.push(c);
264        } else if started {
265            break;
266        }
267    }
268    if buf.is_empty() {
269        return None;
270    }
271    buf.parse::<f64>().ok().or_else(|| {
272        let int_part: String = buf.chars().take_while(|c| c.is_ascii_digit()).collect();
273        int_part.parse::<f64>().ok()
274    })
275}
276
277#[cfg(test)]
278mod tests {
279    use super::*;
280    use lc_embeddings::{EmbeddingError, MockEmbeddings};
281
282    /// 让 embed_query 按文本返回不同维度的向量,复现"向量长度不匹配"。
283    struct MismatchedDimEmbeddings;
284
285    #[async_trait]
286    impl Embeddings for MismatchedDimEmbeddings {
287        async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
288            if text == "pred" {
289                Ok(vec![1.0; 4])
290            } else {
291                Ok(vec![1.0; 8])
292            }
293        }
294        fn dimension(&self) -> usize {
295            4
296        }
297        fn model_name(&self) -> &str {
298            "mismatched-dim"
299        }
300    }
301
302    #[tokio::test]
303    async fn test_exact_match() {
304        let ev = ExactMatch;
305        assert_eq!(ev.eval("", "hello", "hello").await.unwrap().value, 1.0);
306        assert_eq!(ev.eval("", "hello", "world").await.unwrap().value, 0.0);
307        assert_eq!(ev.eval("", "  yes  ", "yes").await.unwrap().value, 1.0);
308    }
309
310    #[tokio::test]
311    async fn test_string_distance() {
312        let ev = StringDistance;
313        let s = ev.eval("", "hello", "hello").await.unwrap().value;
314        assert!((s - 1.0).abs() < 1e-9);
315        let s = ev.eval("", "kitten", "sitting").await.unwrap().value;
316        assert!((s - (1.0 - 3.0 / 7.0)).abs() < 1e-6);
317    }
318
319    #[tokio::test]
320    async fn test_embedding_similarity_identical() {
321        let ev = EmbeddingSimilarity::new(MockEmbeddings::new(32));
322        let s = ev.eval("", "hello", "hello").await.unwrap().value;
323        assert!((s - 1.0).abs() < 1e-6);
324    }
325
326    /// P2-7: 向量长度不匹配是数据缺陷,报错而不是静默返回 0.0。
327    #[tokio::test]
328    async fn test_embedding_similarity_mismatched_dim_errors() {
329        let ev = EmbeddingSimilarity::new(MismatchedDimEmbeddings);
330        let err = ev.eval("", "pred", "ref").await.unwrap_err();
331        assert!(matches!(err, EvalError::EmbeddingError(_)));
332    }
333
334    #[test]
335    fn test_parse_score_unit() {
336        assert!((parse_score(r#"{"score":10}"#, 10).unwrap() - 1.0).abs() < 1e-9);
337        assert!((parse_score(r#"{"score":7.5}"#, 10).unwrap() - 0.75).abs() < 1e-9);
338        assert!((parse_score(r#"{"score":12}"#, 10).unwrap() - 1.0).abs() < 1e-9);
339        assert!(parse_score("no number here", 10).is_none());
340    }
341}