1use async_trait::async_trait;
19use futures_util::stream::{self, StreamExt};
20use serde::Deserialize;
21
22use lc_core::judge::{structured_call, truncate, StructuredJudgeError};
23use lc_core::tools::ToolDefinition;
24use lc_core::BaseChatModel;
25use lc_embeddings::{cosine_similarity, Embeddings};
26use lc_schema::Message;
27
28use super::criteria::{EvalError, Evaluator, RagEvaluator, Score};
29use super::faithfulness::{parse_yes_no, split_claims};
30
31const MAX_CONCURRENT_JUDGE: usize = 4;
34
35const DEFAULT_MAX_CONTEXT_CHARS: usize = 2000;
37
38const DEFAULT_N_QUESTIONS: usize = 3;
40
41#[derive(Debug, Deserialize)]
43struct RagVerdictArgs {
44 verdict: bool,
45 #[serde(default)]
47 #[allow(dead_code)]
48 reason: String,
49}
50
51fn parse_verdict_or_error(raw: &str) -> Result<RagVerdictArgs, StructuredJudgeError> {
54 let verdict = parse_yes_no(raw).ok_or_else(|| {
55 StructuredJudgeError::Parse(format!(
56 "failed to parse yes/no from judge reply: {}",
57 truncate(raw, 200)
58 ))
59 })?;
60 Ok(RagVerdictArgs {
61 verdict,
62 reason: String::new(),
63 })
64}
65
66pub struct ContextPrecision<M: BaseChatModel> {
88 judge: M,
89 max_context_chars: usize,
91 empty_score: f64,
93}
94
95impl<M: BaseChatModel> ContextPrecision<M> {
96 pub fn new(judge: M) -> Self {
98 Self {
99 judge,
100 max_context_chars: DEFAULT_MAX_CONTEXT_CHARS,
101 empty_score: 0.0,
102 }
103 }
104
105 pub fn with_max_context_chars(mut self, max: usize) -> Self {
107 self.max_context_chars = max;
108 self
109 }
110
111 pub fn with_empty_score(mut self, score: f64) -> Self {
113 self.empty_score = score;
114 self
115 }
116
117 async fn judge_chunk(&self, input: &str, chunk: &str) -> Result<bool, EvalError> {
119 let system = "你是检索质量评估员。判断给定的检索文本块是否包含有助于回答用户问题的信息。调用 judge_context 工具提交判定。"
120 .to_string();
121 let user =
122 format!("用户问题:\n{input}\n\n检索文本块:\n{chunk}\n\n该文本块与回答该问题相关吗?");
123 let messages = vec![Message::system(system), Message::human(user)];
124 let args: RagVerdictArgs = structured_call(
125 &self.judge,
126 relevance_tool(),
127 messages,
128 parse_verdict_or_error,
129 )
130 .await?;
131 Ok(args.verdict)
132 }
133}
134
135#[async_trait]
136impl<M: BaseChatModel> RagEvaluator for ContextPrecision<M> {
137 async fn eval_rag(
138 &self,
139 input: &str,
140 _prediction: &str,
141 contexts: &[String],
142 _reference: &str,
143 ) -> Result<Score, EvalError> {
144 if contexts.is_empty() {
145 return Ok(Score::new(self.empty_score).with_label("no_contexts"));
146 }
147 let chunks: Vec<String> = contexts
150 .iter()
151 .map(|c| truncate(c, self.max_context_chars).to_string())
152 .collect();
153 let verdicts: Vec<Result<bool, EvalError>> = stream::iter(chunks)
154 .map(|chunk| async move { self.judge_chunk(input, &chunk).await })
155 .buffered(MAX_CONCURRENT_JUDGE)
156 .collect()
157 .await;
158
159 let mut relevant_in_top_k = 0usize;
160 let mut total_relevant = 0usize;
161 let mut weighted = 0.0;
162 for (k, verdict) in verdicts.into_iter().enumerate() {
163 let relevant = verdict?;
164 if relevant {
165 relevant_in_top_k += 1;
166 total_relevant += 1;
167 let precision_at_k = relevant_in_top_k as f64 / (k + 1) as f64;
168 weighted += precision_at_k;
169 }
170 }
171 if total_relevant == 0 {
172 return Ok(Score::new(0.0).with_label("no_relevant"));
174 }
175 Ok(Score::new(weighted / total_relevant as f64).with_label("context_precision"))
176 }
177
178 fn name(&self) -> &str {
179 "context_precision"
180 }
181}
182
183fn relevance_tool() -> ToolDefinition {
184 ToolDefinition::new(
185 "judge_context",
186 "判断检索文本块是否与用户问题相关,提交布尔判定。",
187 )
188 .with_parameters(serde_json::json!({
189 "type": "object",
190 "properties": {
191 "verdict": { "type": "boolean", "description": "文本块是否包含有助于回答问题的信息" },
192 "reason": { "type": "string", "description": "简短依据" }
193 },
194 "required": ["verdict", "reason"]
195 }))
196}
197
198pub struct ContextRecall<M: BaseChatModel> {
207 judge: M,
208 max_context_chars: usize,
210 empty_score: f64,
212}
213
214impl<M: BaseChatModel> ContextRecall<M> {
215 pub fn new(judge: M) -> Self {
217 Self {
218 judge,
219 max_context_chars: DEFAULT_MAX_CONTEXT_CHARS,
220 empty_score: 0.0,
221 }
222 }
223
224 pub fn with_max_context_chars(mut self, max: usize) -> Self {
226 self.max_context_chars = max;
227 self
228 }
229
230 pub fn with_empty_score(mut self, score: f64) -> Self {
232 self.empty_score = score;
233 self
234 }
235
236 async fn verify_claim(&self, context: &str, claim: &str) -> Result<bool, EvalError> {
238 let system = "你是事实核查员。判断参考答案中的陈述能否从任一检索上下文中推导出来。调用 check_claim 工具提交判定。"
239 .to_string();
240 let user = format!(
241 "检索上下文:\n{context}\n\n参考答案陈述:\n{claim}\n\n这条陈述能从检索上下文推导出来吗?"
242 );
243 let messages = vec![Message::system(system), Message::human(user)];
244 let args: RagVerdictArgs =
245 structured_call(&self.judge, recall_tool(), messages, parse_verdict_or_error).await?;
246 Ok(args.verdict)
247 }
248}
249
250#[async_trait]
251impl<M: BaseChatModel> RagEvaluator for ContextRecall<M> {
252 async fn eval_rag(
253 &self,
254 _input: &str,
255 _prediction: &str,
256 contexts: &[String],
257 reference: &str,
258 ) -> Result<Score, EvalError> {
259 if contexts.is_empty() {
260 return Ok(Score::new(self.empty_score).with_label("no_contexts"));
261 }
262 let claims = split_claims(reference);
263 if claims.is_empty() {
264 return Ok(Score::new(self.empty_score).with_label("no_claims"));
265 }
266 let context = truncate(&contexts.join("\n\n---\n\n"), self.max_context_chars);
268 let ctx = &context;
269 let total = claims.len();
270 let results: Vec<Result<bool, EvalError>> = stream::iter(claims)
271 .map(|claim| async move { self.verify_claim(ctx, &claim).await })
272 .buffer_unordered(MAX_CONCURRENT_JUDGE)
273 .collect()
274 .await;
275 let mut attributable = 0usize;
276 for r in results {
277 if r? {
278 attributable += 1;
279 }
280 }
281 Ok(Score::new(attributable as f64 / total as f64).with_label("context_recall"))
282 }
283
284 fn name(&self) -> &str {
285 "context_recall"
286 }
287}
288
289fn recall_tool() -> ToolDefinition {
290 ToolDefinition::new(
291 "check_claim",
292 "判断参考答案陈述能否从检索上下文推导出来,提交布尔判定。",
293 )
294 .with_parameters(serde_json::json!({
295 "type": "object",
296 "properties": {
297 "verdict": { "type": "boolean", "description": "能否从任一检索上下文推导" },
298 "reason": { "type": "string", "description": "简短依据" }
299 },
300 "required": ["verdict", "reason"]
301 }))
302}
303
304pub struct AnswerRelevancy<M: BaseChatModel, E: Embeddings> {
317 generator: M,
318 embeddings: E,
319 n_questions: usize,
321 empty_score: f64,
323}
324
325impl<M: BaseChatModel, E: Embeddings> AnswerRelevancy<M, E> {
326 pub fn new(generator: M, embeddings: E) -> Self {
328 Self {
329 generator,
330 embeddings,
331 n_questions: DEFAULT_N_QUESTIONS,
332 empty_score: 0.0,
333 }
334 }
335
336 pub fn with_n_questions(mut self, n: usize) -> Self {
338 self.n_questions = n.max(1);
339 self
340 }
341
342 pub fn with_empty_score(mut self, score: f64) -> Self {
344 self.empty_score = score;
345 self
346 }
347
348 async fn score(&self, input: &str, prediction: &str) -> Result<Score, EvalError> {
350 if prediction.trim().is_empty() {
351 return Ok(Score::new(self.empty_score).with_label("no_answer"));
352 }
353 let questions = self.generate_questions(prediction).await?;
354 if questions.is_empty() {
355 return Err(EvalError::ParseError(
357 "answer relevancy generator produced no questions".into(),
358 ));
359 }
360
361 let original = self
362 .embeddings
363 .embed_query(input)
364 .await
365 .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
366 let refs: Vec<&str> = questions.iter().map(String::as_str).collect();
367 let generated = self
368 .embeddings
369 .embed_documents(&refs)
370 .await
371 .map_err(|e| EvalError::EmbeddingError(e.to_string()))?;
372 if generated.len() != questions.len() {
373 return Err(EvalError::EmbeddingError(format!(
374 "embedding batch mismatch: asked for {}, got {}",
375 questions.len(),
376 generated.len()
377 )));
378 }
379
380 let mut sum = 0.0;
381 for v in &generated {
382 let sim = cosine_similarity(&original, v)
384 .map_err(|e| EvalError::EmbeddingError(e.to_string()))?
385 as f64;
386 sum += sim;
387 }
388 Ok(Score::new(sum / questions.len() as f64).with_label("answer_relevancy"))
390 }
391
392 async fn generate_questions(&self, prediction: &str) -> Result<Vec<String>, EvalError> {
394 let system = format!(
395 "你是问题生成器。仅根据给定回答,生成 {} 个不同的、该回答能够回答的问题。每行一个问题,不要编号、不要解释。",
396 self.n_questions
397 );
398 let user = format!(
399 "回答:\n{prediction}\n\n请生成 {} 个问题,每行一个:",
400 self.n_questions
401 );
402 let result = self
403 .generator
404 .chat_with_system(system, vec![Message::human(user)])
405 .await
406 .map_err(|e| EvalError::PredictorError(e.to_string()))?;
407 Ok(result
408 .content
409 .lines()
410 .map(str::trim)
411 .map(|l| {
412 let bytes = l.as_bytes();
415 let mut i = 0;
416 while i < bytes.len() && bytes[i].is_ascii_digit() {
417 i += 1;
418 }
419 let ascii_sep = i < bytes.len() && (bytes[i] == b'.' || bytes[i] == b')');
420 let ideographic_sep = i < bytes.len() && l[i..].starts_with('、');
421 if i > 0 && (ascii_sep || ideographic_sep) {
422 let sep_len = if ascii_sep { 1 } else { '、'.len_utf8() };
423 l[i + sep_len..].trim()
424 } else {
425 l
426 }
427 })
428 .filter(|l| !l.is_empty())
429 .map(str::to_string)
430 .collect())
431 }
432}
433
434#[async_trait]
435impl<M: BaseChatModel, E: Embeddings> RagEvaluator for AnswerRelevancy<M, E> {
436 async fn eval_rag(
437 &self,
438 input: &str,
439 prediction: &str,
440 _contexts: &[String],
441 _reference: &str,
442 ) -> Result<Score, EvalError> {
443 self.score(input, prediction).await
444 }
445
446 fn name(&self) -> &str {
447 "answer_relevancy"
448 }
449}
450
451#[async_trait]
452impl<M: BaseChatModel, E: Embeddings> Evaluator for AnswerRelevancy<M, E> {
453 async fn eval(
454 &self,
455 input: &str,
456 prediction: &str,
457 _reference: &str,
458 ) -> Result<Score, EvalError> {
459 self.score(input, prediction).await
460 }
461
462 fn name(&self) -> &str {
463 "answer_relevancy"
464 }
465}
466
467#[cfg(test)]
468mod tests {
469 use super::*;
470 use futures_util::Stream;
471 use lc_core::language_models::{LLMResult, StreamChunk};
472 use lc_core::{BaseLanguageModel, Runnable, RunnableConfig};
473 use lc_embeddings::EmbeddingError;
474 use std::pin::Pin;
475 use std::sync::atomic::{AtomicUsize, Ordering};
476 use std::sync::Arc;
477
478 #[derive(Debug)]
479 struct MockError(String);
480 impl std::fmt::Display for MockError {
481 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
482 write!(f, "{}", self.0)
483 }
484 }
485 impl std::error::Error for MockError {}
486
487 struct TextMock {
489 replies: Vec<String>,
490 calls: Arc<AtomicUsize>,
491 }
492 impl TextMock {
493 fn new(replies: Vec<String>) -> Self {
494 Self {
495 replies,
496 calls: Arc::new(AtomicUsize::new(0)),
497 }
498 }
499 }
500
501 #[async_trait]
502 impl Runnable<Vec<Message>, LLMResult> for TextMock {
503 type Error = MockError;
504 async fn invoke(
505 &self,
506 _input: Vec<Message>,
507 _config: Option<RunnableConfig>,
508 ) -> Result<LLMResult, Self::Error> {
509 Err(MockError("use chat".into()))
510 }
511 }
512 #[async_trait]
513 impl BaseLanguageModel<Vec<Message>, LLMResult> for TextMock {
514 fn model_name(&self) -> &str {
515 "text-mock"
516 }
517 fn get_num_tokens(&self, t: &str) -> usize {
518 t.len()
519 }
520 fn with_temperature(self, _: f32) -> Self {
521 self
522 }
523 fn with_max_tokens(self, _: usize) -> Self {
524 self
525 }
526 }
527 #[async_trait]
528 impl BaseChatModel for TextMock {
529 async fn chat(
530 &self,
531 _messages: Vec<Message>,
532 _config: Option<RunnableConfig>,
533 ) -> Result<LLMResult, Self::Error> {
534 let idx = self.calls.fetch_add(1, Ordering::SeqCst);
535 Ok(LLMResult {
536 content: self.replies.get(idx).cloned().unwrap_or_default(),
537 model: "text-mock".into(),
538 token_usage: None,
539 tool_calls: None,
540 thinking_content: None,
541 })
542 }
543 async fn stream_chat(
544 &self,
545 _messages: Vec<Message>,
546 _config: Option<RunnableConfig>,
547 ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamChunk, Self::Error>> + Send>>, Self::Error>
548 {
549 Err(MockError("not supported".into()))
550 }
551 }
552
553 use crate::test_support::ToolJudge;
555
556 struct ScriptedEmbeddings {
558 dim: usize,
559 map: Vec<(String, Vec<f32>)>,
560 }
561 impl ScriptedEmbeddings {
562 fn new(map: Vec<(&str, Vec<f32>)>) -> Self {
563 let dim = map.first().map(|(_, v)| v.len()).unwrap_or(1);
564 Self {
565 dim,
566 map: map.into_iter().map(|(k, v)| (k.to_string(), v)).collect(),
567 }
568 }
569 }
570 #[async_trait]
571 impl Embeddings for ScriptedEmbeddings {
572 async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
573 self.map
574 .iter()
575 .find(|(k, _)| k == text)
576 .map(|(_, v)| v.clone())
577 .ok_or_else(|| EmbeddingError::Config(format!("unscripted text: {text}")))
578 }
579 fn dimension(&self) -> usize {
580 self.dim
581 }
582 fn model_name(&self) -> &str {
583 "scripted"
584 }
585 }
586
587 #[tokio::test]
590 async fn context_precision_weights_by_rank() {
591 let judge = ToolJudge::sequence(vec![
594 r#"{"verdict": true, "reason": "r"}"#.into(),
595 r#"{"verdict": false, "reason": "r"}"#.into(),
596 r#"{"verdict": true, "reason": "r"}"#.into(),
597 ]);
598 let contexts = vec!["c0".into(), "c1".into(), "c2".into()];
599 let s = ContextPrecision::new(judge)
600 .eval_rag("q", "a", &contexts, "ref")
601 .await
602 .unwrap();
603 assert!((s.value - 5.0 / 6.0).abs() < 1e-9, "got {}", s.value);
604 }
605
606 #[tokio::test]
607 async fn context_precision_all_relevant_is_one() {
608 let judge = ToolJudge::sequence(vec![
609 r#"{"verdict": true}"#.into(),
610 r#"{"verdict": true}"#.into(),
611 ]);
612 let contexts = vec!["c0".into(), "c1".into()];
613 let s = ContextPrecision::new(judge)
614 .eval_rag("q", "a", &contexts, "ref")
615 .await
616 .unwrap();
617 assert!((s.value - 1.0).abs() < 1e-9);
618 }
619
620 #[tokio::test]
621 async fn context_precision_none_relevant_is_zero() {
622 let judge = ToolJudge::sequence(vec![
623 r#"{"verdict": false}"#.into(),
624 r#"{"verdict": false}"#.into(),
625 ]);
626 let contexts = vec!["c0".into(), "c1".into()];
627 let s = ContextPrecision::new(judge)
628 .eval_rag("q", "a", &contexts, "ref")
629 .await
630 .unwrap();
631 assert_eq!(s.value, 0.0);
632 assert_eq!(s.label.as_deref(), Some("no_relevant"));
633 }
634
635 #[tokio::test]
636 async fn context_precision_empty_contexts_uses_empty_score() {
637 let judge = ToolJudge::new(r#"{"verdict": true}"#);
638 let s = ContextPrecision::new(judge)
639 .with_empty_score(1.0)
640 .eval_rag("q", "a", &[], "ref")
641 .await
642 .unwrap();
643 assert_eq!(s.value, 1.0);
644 assert_eq!(s.label.as_deref(), Some("no_contexts"));
645 }
646
647 #[tokio::test]
650 async fn context_recall_half_attributable() {
651 let judge = ToolJudge::sequence(vec![
652 r#"{"verdict": true, "reason": "r"}"#.into(),
653 r#"{"verdict": false, "reason": "r"}"#.into(),
654 ]);
655 let contexts = vec!["ctx".into()];
656 let s = ContextRecall::new(judge)
657 .eval_rag("q", "a", &contexts, "巴黎是首都。伦敦是首都。")
658 .await
659 .unwrap();
660 assert!((s.value - 0.5).abs() < 1e-9);
661 }
662
663 #[tokio::test]
664 async fn context_recall_no_contexts() {
665 let judge = ToolJudge::new(r#"{"verdict": true}"#);
666 let s = ContextRecall::new(judge)
667 .eval_rag("q", "a", &[], "巴黎是首都。")
668 .await
669 .unwrap();
670 assert_eq!(s.value, 0.0);
671 assert_eq!(s.label.as_deref(), Some("no_contexts"));
672 }
673
674 #[tokio::test]
675 async fn context_recall_no_claims() {
676 let judge = ToolJudge::new(r#"{"verdict": true}"#);
677 let s = ContextRecall::new(judge)
678 .with_empty_score(1.0)
679 .eval_rag("q", "a", &["ctx".to_string()], "。。。")
680 .await
681 .unwrap();
682 assert_eq!(s.value, 1.0);
683 assert_eq!(s.label.as_deref(), Some("no_claims"));
684 }
685
686 #[tokio::test]
689 async fn answer_relevancy_identical_questions_scores_one() {
690 let gen = TextMock::new(vec!["q-gen-1\nq-gen-2".into()]);
691 let emb = ScriptedEmbeddings::new(vec![
692 ("q", vec![1.0, 0.0]),
693 ("q-gen-1", vec![1.0, 0.0]),
694 ("q-gen-2", vec![1.0, 0.0]),
695 ]);
696 let s = AnswerRelevancy::new(gen, emb)
697 .eval_rag("q", "an answer", &[], "")
698 .await
699 .unwrap();
700 assert!((s.value - 1.0).abs() < 1e-6);
701 assert_eq!(s.label.as_deref(), Some("answer_relevancy"));
702 }
703
704 #[tokio::test]
705 async fn answer_relevancy_averages_cosines() {
706 let gen = TextMock::new(vec!["same\northogonal".into()]);
707 let emb = ScriptedEmbeddings::new(vec![
708 ("q", vec![1.0, 0.0]),
709 ("same", vec![1.0, 0.0]),
710 ("orthogonal", vec![0.0, 1.0]),
711 ]);
712 let s = AnswerRelevancy::new(gen, emb)
713 .eval("q", "an answer", "")
714 .await
715 .unwrap();
716 assert!((s.value - 0.5).abs() < 1e-6, "got {}", s.value);
717 }
718
719 #[tokio::test]
720 async fn answer_relevancy_strips_list_numbering() {
721 let gen = TextMock::new(vec!["1. same\n2) same".into()]);
723 let emb = ScriptedEmbeddings::new(vec![("q", vec![1.0, 0.0]), ("same", vec![1.0, 0.0])]);
724 let s = AnswerRelevancy::new(gen, emb)
725 .eval("q", "answer", "")
726 .await
727 .unwrap();
728 assert!((s.value - 1.0).abs() < 1e-6);
729 }
730
731 #[tokio::test]
732 async fn answer_relevancy_keeps_leading_digits_of_real_questions() {
733 let gen = TextMock::new(vec!["2+2等于几?".into()]);
735 let emb =
736 ScriptedEmbeddings::new(vec![("q", vec![1.0, 0.0]), ("2+2等于几?", vec![1.0, 0.0])]);
737 let s = AnswerRelevancy::new(gen, emb)
738 .eval("q", "answer", "")
739 .await
740 .unwrap();
741 assert!((s.value - 1.0).abs() < 1e-6);
742 }
743
744 #[tokio::test]
745 async fn answer_relevancy_empty_prediction_uses_empty_score() {
746 let gen = TextMock::new(vec![]);
747 let emb = ScriptedEmbeddings::new(vec![("q", vec![1.0])]);
748 let s = AnswerRelevancy::new(gen, emb)
749 .with_empty_score(1.0)
750 .eval("q", " ", "")
751 .await
752 .unwrap();
753 assert_eq!(s.value, 1.0);
754 assert_eq!(s.label.as_deref(), Some("no_answer"));
755 }
756
757 #[tokio::test]
758 async fn answer_relevancy_zero_generated_questions_errors() {
759 let gen = TextMock::new(vec![" \n ".into()]);
760 let emb = ScriptedEmbeddings::new(vec![("q", vec![1.0])]);
761 let err = AnswerRelevancy::new(gen, emb)
762 .eval("q", "answer", "")
763 .await
764 .unwrap_err();
765 assert!(matches!(err, EvalError::ParseError(_)), "got {err:?}");
766 }
767
768 #[tokio::test]
769 async fn answer_relevancy_dimension_mismatch_errors() {
770 let gen = TextMock::new(vec!["g".into()]);
771 let emb = ScriptedEmbeddings::new(vec![("q", vec![1.0, 0.0]), ("g", vec![1.0, 0.0, 0.0])]);
772 let err = AnswerRelevancy::new(gen, emb)
773 .eval("q", "answer", "")
774 .await
775 .unwrap_err();
776 assert!(matches!(err, EvalError::EmbeddingError(_)));
777 }
778}