diff --git a/README.md b/README.md index 73177a5..984a2b7 100644 --- a/README.md +++ b/README.md @@ -140,7 +140,30 @@ uv run local-ai-agent evaluate --report-dir evaluation/results/my-run The versioned case manifest is tied to the dataset SHA-256 and uses immutable content-derived source IDs for gold relevance. The report records model tags and immutable Ollama digests, dataset and case-set hashes, retrieval limit, -runtime versions, per-case RAG and BM25 rankings, and aggregate metrics. +runtime versions, per-case RAG and BM25 rankings, and aggregate metrics. Each +RAG observation also retains the first raw model response, any one-shot repair +response, the initial and final structured validation reasons, and whether a +repair was attempted. These diagnostics stay in the evaluation artifact; the +CLI and dashboard continue to expose only validated answers or safe fallback +messages. + +Compare answer models without changing the embedding model or case set: + +```bash +uv run local-ai-agent evaluate \ + --chat-model llama3.2 \ + --report-dir /tmp/local-ai-agent-llama3.2 + +uv run local-ai-agent evaluate \ + --chat-model \ + --report-dir /tmp/local-ai-agent-stronger-model +``` + +Keep the reports outside the repository so benchmark artifacts do not make the +worktree dirty. Compare answer success, citation validity, abstention recall, +latency, and the per-case failure reasons rather than optimizing one aggregate +score. + Retrieval quality is reported as recall@k, hit rate@k, and MRR@k for both semantic search and the BM25 baseline. Relevance judgments are known-positive, not exhaustive. The generated report is evidence for this fixed benchmark diff --git a/agent.py b/agent.py index be67dfc..d291b30 100644 --- a/agent.py +++ b/agent.py @@ -17,13 +17,19 @@ ) CITATION_PATTERN = re.compile(r"\[([A-Za-z0-9][A-Za-z0-9_-]*)\]") INSUFFICIENT_EVIDENCE_TOKEN = "INSUFFICIENT_EVIDENCE" +REPAIRABLE_FAILURE_REASONS = frozenset( + {"missing_citations", "out_of_range_citation", "unknown_citation"} +) ANSWER_PROMPT = """You are a review analyst. Answer the question using only the supplied reviews. Do not add facts that are not present. -When evidence is mixed or limited, say so clearly. Every factual claim must cite one or more -retrieved evidence numbers exactly as shown, for example [1]. Never cite an evidence number -that is not supplied. Source IDs are validation metadata; do not copy them into the answer. -If the supplied reviews do not answer the question, reply exactly INSUFFICIENT_EVIDENCE. +Always answer every part that at least one supplied review supports. Mixed or incomplete evidence is +not a reason to abstain; describe the limitation and answer only the supported portion. Every +factual claim must cite one or more retrieved evidence numbers exactly as shown, for example +[1]. Never cite an evidence number that is not supplied. Source IDs are validation metadata; +do not copy them into the answer. Reply exactly INSUFFICIENT_EVIDENCE only when no supplied +review answers any part of the question. Never use INSUFFICIENT_EVIDENCE as prose or place it +beside an answer. Question: {question} @@ -34,6 +40,27 @@ Answer: """ +REPAIR_PROMPT = """You are repairing one rejected review-analysis response. +Using only the supplied reviews, rewrite it once so every factual claim has one or more valid +evidence citations such as [1]. Use only evidence numbers that appear below. Preserve supported +meaning, remove unsupported claims, and answer every supported part even when evidence is mixed +or incomplete. If no supplied review answers any part, reply exactly INSUFFICIENT_EVIDENCE. +Never use INSUFFICIENT_EVIDENCE as prose or place it beside an answer. + +Question: +{question} + +Supplied review records: +{context} + +Rejected response: + +{rejected_response} + + +Rewritten answer: +""" + @dataclass(frozen=True) class CitedReview: @@ -48,6 +75,11 @@ class AnswerResult: sources: tuple[CitedReview, ...] retrieved_source_ids: tuple[str, ...] = () abstained: bool = False + raw_response: str = "" + repair_response: str | None = None + initial_failure_reason: str | None = None + failure_reason: str | None = None + repair_attempted: bool = False def create_chat_model( @@ -99,7 +131,7 @@ def _remove_standalone_control_token(answer: str) -> str | None: def _validate_and_number_citations( answer: str, matches: list[ReviewMatch], -) -> tuple[str, tuple[CitedReview, ...]] | None: +) -> tuple[tuple[str, tuple[CitedReview, ...]] | None, str | None]: retrieved: dict[str, ReviewMatch] = {} evidence_aliases: dict[str, str] = {} for evidence_number, match in enumerate(matches, start=1): @@ -107,19 +139,20 @@ def _validate_and_number_citations( match.document.metadata.get("source_id") or match.document.id or "" ) if not source_id: - return None + return None, "retrieved_source_missing_id" retrieved[source_id] = match evidence_aliases[str(evidence_number)] = source_id cited_tokens = CITATION_PATTERN.findall(answer) if not cited_tokens: - return None + return None, "missing_citations" resolved_ids: list[str] = [] for token in cited_tokens: source_id = evidence_aliases.get(token, token) if source_id not in retrieved: - return None + reason = "out_of_range_citation" if token.isdigit() else "unknown_citation" + return None, reason resolved_ids.append(source_id) ordered_ids = list(dict.fromkeys(resolved_ids)) @@ -140,7 +173,26 @@ def _validate_and_number_citations( ) for source_id in ordered_ids ) - return numbered_answer, sources + return (numbered_answer, sources), None + + +def _evaluate_model_response( + response: str, + matches: list[ReviewMatch], +) -> tuple[tuple[str, tuple[CitedReview, ...]] | None, str | None]: + normalized = response.strip() + if normalized == INSUFFICIENT_EVIDENCE_TOKEN: + return None, "clean_abstention" + normalized_without_control = _remove_standalone_control_token(normalized) + if normalized_without_control is None: + return None, "embedded_control_token" + if not normalized_without_control: + return None, "invalid_remainder" + return _validate_and_number_citations(normalized_without_control, matches) + + +def _response_text(response: Any) -> str: + return str(response.content if hasattr(response, "content") else response) def answer_question( @@ -177,47 +229,95 @@ def answer_question( countries=countries, ) if not matches: - return AnswerResult(answer=NO_MATCH_MESSAGE, sources=()) + return AnswerResult( + answer=NO_MATCH_MESSAGE, + sources=(), + failure_reason="empty_retrieval", + ) retrieved_source_ids = tuple( str(match.document.metadata.get("source_id") or match.document.id or "") for match in matches ) if any(not source_id for source_id in retrieved_source_ids): - return AnswerResult(answer=CITATION_VALIDATION_MESSAGE, sources=()) + return AnswerResult( + answer=CITATION_VALIDATION_MESSAGE, + sources=(), + failure_reason="retrieved_source_missing_id", + ) answer_model = model or create_chat_model(model=chat_model, base_url=ollama_host) + context = _format_context(matches) prompt = ANSWER_PROMPT.format( question=normalized_question, - context=_format_context(matches), + context=context, ) - response = answer_model.invoke(prompt) - answer = response.content if hasattr(response, "content") else str(response) - normalized_answer = answer.strip() - if normalized_answer == INSUFFICIENT_EVIDENCE_TOKEN: + raw_response = _response_text(answer_model.invoke(prompt)) + validated, failure_reason = _evaluate_model_response(raw_response, matches) + if validated is not None: + validated_answer, sources = validated + return AnswerResult( + answer=validated_answer, + sources=sources, + retrieved_source_ids=retrieved_source_ids, + raw_response=raw_response, + ) + if failure_reason == "clean_abstention": return AnswerResult( answer=NO_MATCH_MESSAGE, sources=(), retrieved_source_ids=retrieved_source_ids, abstained=True, + raw_response=raw_response, + failure_reason=failure_reason, ) - normalized_answer = _remove_standalone_control_token(normalized_answer) - if normalized_answer is None: + if failure_reason not in REPAIRABLE_FAILURE_REASONS: return AnswerResult( answer=CITATION_VALIDATION_MESSAGE, sources=(), retrieved_source_ids=retrieved_source_ids, + raw_response=raw_response, + failure_reason=failure_reason, ) - validated = _validate_and_number_citations(normalized_answer, matches) - if validated is None: + + initial_failure_reason = failure_reason + repair_prompt = REPAIR_PROMPT.format( + question=normalized_question, + context=context, + rejected_response=raw_response, + ) + repair_response = _response_text(answer_model.invoke(repair_prompt)) + repaired, repair_failure_reason = _evaluate_model_response(repair_response, matches) + if repaired is not None: + repaired_answer, sources = repaired return AnswerResult( - answer=CITATION_VALIDATION_MESSAGE, + answer=repaired_answer, + sources=sources, + retrieved_source_ids=retrieved_source_ids, + raw_response=raw_response, + repair_response=repair_response, + initial_failure_reason=initial_failure_reason, + repair_attempted=True, + ) + if repair_failure_reason == "clean_abstention": + return AnswerResult( + answer=NO_MATCH_MESSAGE, sources=(), retrieved_source_ids=retrieved_source_ids, + abstained=True, + raw_response=raw_response, + repair_response=repair_response, + initial_failure_reason=initial_failure_reason, + failure_reason=repair_failure_reason, + repair_attempted=True, ) - validated_answer, sources = validated return AnswerResult( - answer=validated_answer, - sources=sources, + answer=CITATION_VALIDATION_MESSAGE, + sources=(), retrieved_source_ids=retrieved_source_ids, + raw_response=raw_response, + repair_response=repair_response, + initial_failure_reason=initial_failure_reason, + failure_reason=repair_failure_reason, + repair_attempted=True, ) diff --git a/evaluation.py b/evaluation.py index 3fec482..4146cf0 100644 --- a/evaluation.py +++ b/evaluation.py @@ -50,6 +50,11 @@ class EvaluationObservation: abstained: bool outcome: str = "unknown" latency_ms: float | None = None + raw_model_response: str = "" + repair_model_response: str | None = None + initial_failure_reason: str | None = None + failure_reason: str | None = None + repair_attempted: bool = False @dataclass(frozen=True) @@ -409,7 +414,10 @@ def score_evaluation( cited = set(observation.cited_source_ids) retrieved = set(observation.retrieved_source_ids) - answer_succeeded = observation.outcome == "answered" or ( + answer_succeeded = observation.outcome in { + "answered", + "answered_after_repair", + } or ( observation.outcome == "unknown" and not observation.abstained and bool(cited) @@ -472,9 +480,19 @@ def run_rag_evaluation( if not result.retrieved_source_ids: outcome = "empty_retrieval" elif result.abstained: - outcome = "model_abstention" + outcome = ( + "model_abstention_after_repair" + if result.repair_attempted + else "model_abstention" + ) elif not result.sources: - outcome = "citation_validation_rejection" + outcome = ( + "citation_validation_rejection_after_repair" + if result.repair_attempted + else "citation_validation_rejection" + ) + elif result.repair_attempted: + outcome = "answered_after_repair" else: outcome = "answered" observations.append( @@ -492,6 +510,11 @@ def run_rag_evaluation( abstained=result.abstained, outcome=outcome, latency_ms=latency_ms, + raw_model_response=result.raw_response, + repair_model_response=result.repair_response, + initial_failure_reason=result.initial_failure_reason, + failure_reason=result.failure_reason, + repair_attempted=result.repair_attempted, ) ) @@ -626,8 +649,19 @@ def build_evaluation_report( for observation in observations if observation.latency_ms is not None ] + outcome_counts = Counter(observation.outcome for observation in observations) + initial_failure_counts = Counter( + observation.initial_failure_reason + for observation in observations + if observation.initial_failure_reason is not None + ) + final_failure_counts = Counter( + observation.failure_reason + for observation in observations + if observation.failure_reason is not None + ) return { - "schema_version": 2, + "schema_version": 3, "generated_at": timestamp, "configuration": dict(configuration), "provenance": dict(provenance), @@ -647,6 +681,15 @@ def build_evaluation_report( "rag_mean_latency_ms": round(_mean(latencies), 3), "rag_median_latency_ms": round(median(latencies), 3) if latencies else 0.0, }, + "diagnostics": { + "outcomes": dict(sorted(outcome_counts.items())), + "initial_failure_reasons": dict(sorted(initial_failure_counts.items())), + "final_failure_reasons": dict(sorted(final_failure_counts.items())), + "repair_attempt_count": sum( + observation.repair_attempted for observation in observations + ), + "repair_success_count": outcome_counts["answered_after_repair"], + }, "observations": { "rag": [_observation_as_dict(observation) for observation in observations], "bm25_baseline": [ @@ -667,6 +710,11 @@ def _observation_as_dict(observation: EvaluationObservation) -> dict[str, Any]: "abstained": observation.abstained, "outcome": observation.outcome, "latency_ms": observation.latency_ms, + "raw_model_response": observation.raw_model_response, + "repair_model_response": observation.repair_model_response, + "initial_failure_reason": observation.initial_failure_reason, + "failure_reason": observation.failure_reason, + "repair_attempted": observation.repair_attempted, } @@ -685,6 +733,7 @@ def _report_markdown(report: Mapping[str, Any]) -> str: configuration = report["configuration"] provenance = report["provenance"] timing = report["timing"] + diagnostics = report["diagnostics"] return "\n".join( ( "# Evaluation report", @@ -733,6 +782,20 @@ def _report_markdown(report: Mapping[str, Any]) -> str: f"| Answer success (answerable cases) | {_metric(rag['answer_success_rate'])} |", f"| Abstention recall (abstention cases) | {_metric(rag['abstention_recall'])} |", "", + "## Answer diagnostics", + "", + f"- Outcomes: `{json.dumps(diagnostics['outcomes'], sort_keys=True)}`", + ( + "- Initial validation reasons: `" + f"{json.dumps(diagnostics['initial_failure_reasons'], sort_keys=True)}`" + ), + ( + "- Final non-success reasons: `" + f"{json.dumps(diagnostics['final_failure_reasons'], sort_keys=True)}`" + ), + f"- Repair attempts: **{diagnostics['repair_attempt_count']}**", + f"- Successful repairs: **{diagnostics['repair_success_count']}**", + "", "## Timing", "", f"- Total RAG latency: **{timing['rag_total_latency_ms'] / 1000:.2f}s**", diff --git a/tests/test_agent.py b/tests/test_agent.py index a29f859..3a895fa 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -21,6 +21,16 @@ def invoke(self, prompt: str) -> str: return self.response +class SequenceModel: + def __init__(self, *responses: str) -> None: + self.responses = list(responses) + self.prompts: list[str] = [] + + def invoke(self, prompt: str) -> str: + self.prompts.append(prompt) + return self.responses.pop(0) + + class FakeStore: def __init__(self, results: list[tuple[Document, float]]) -> None: self.results = results @@ -59,6 +69,87 @@ def test_builds_grounded_prompt_and_returns_cited_sources(self) -> None: self.assertIn("Source ID: review-1", model.prompts[0]) self.assertIn("Great crust Crisp and flavorful.", model.prompts[0]) self.assertIn("only the supplied reviews", model.prompts[0]) + self.assertIn( + "answer every part that at least one supplied review supports", + model.prompts[0], + ) + self.assertIn("Mixed or incomplete evidence is", model.prompts[0]) + self.assertIn("not a reason to abstain", model.prompts[0]) + self.assertIn("Never use INSUFFICIENT_EVIDENCE as prose", model.prompts[0]) + + def test_repairs_an_uncited_answer_once_and_preserves_diagnostics(self) -> None: + document = Document( + page_content="The crust was perfectly crispy.", + metadata={"source_id": "review-1"}, + id="review-1", + ) + model = SequenceModel( + "Guests praise the crispy crust.", + "Guests praise the crispy crust [1].", + ) + + result = answer_question( + "What do guests say about the crust?", + vector_store=FakeStore([(document, 0.5)]), + model=model, + ) + + self.assertEqual(result.answer, "Guests praise the crispy crust [1].") + self.assertEqual(result.raw_response, "Guests praise the crispy crust.") + self.assertEqual(result.repair_response, "Guests praise the crispy crust [1].") + self.assertEqual(result.initial_failure_reason, "missing_citations") + self.assertIsNone(result.failure_reason) + self.assertTrue(result.repair_attempted) + self.assertEqual(len(model.prompts), 2) + self.assertIn("Guests praise the crispy crust.", model.prompts[1]) + self.assertIn("rewrite it once", model.prompts[1]) + + def test_failed_repair_records_the_final_structured_reason(self) -> None: + document = Document( + page_content="The crust was perfectly crispy.", + metadata={"source_id": "review-1"}, + id="review-1", + ) + model = SequenceModel("Unsupported [2].", "Still unsupported [3].") + + result = answer_question( + "What do guests say about the crust?", + vector_store=FakeStore([(document, 0.5)]), + model=model, + ) + + self.assertEqual(result.answer, CITATION_VALIDATION_MESSAGE) + self.assertEqual(result.initial_failure_reason, "out_of_range_citation") + self.assertEqual(result.failure_reason, "out_of_range_citation") + self.assertEqual(result.raw_response, "Unsupported [2].") + self.assertEqual(result.repair_response, "Still unsupported [3].") + self.assertTrue(result.repair_attempted) + self.assertEqual(len(model.prompts), 2) + + def test_repair_can_return_a_clean_abstention_without_a_third_attempt(self) -> None: + document = Document( + page_content="A review about pizza crust.", + metadata={"source_id": "review-1"}, + id="review-1", + ) + model = SequenceModel( + "Parking is available.", + " INSUFFICIENT_EVIDENCE ", + ) + + result = answer_question( + "Is parking available?", + vector_store=FakeStore([(document, 0.5)]), + model=model, + ) + + self.assertEqual(result.answer, NO_MATCH_MESSAGE) + self.assertEqual(result.sources, ()) + self.assertTrue(result.abstained) + self.assertEqual(result.initial_failure_reason, "missing_citations") + self.assertEqual(result.failure_reason, "clean_abstention") + self.assertTrue(result.repair_attempted) + self.assertEqual(len(model.prompts), 2) def test_accepts_numeric_citation_for_the_matching_retrieved_review(self) -> None: document = Document( @@ -161,16 +252,20 @@ def test_model_can_abstain_when_retrieved_reviews_are_insufficient(self) -> None id="review-1", ) + model = FakeModel("\n INSUFFICIENT_EVIDENCE \n") result = answer_question( "Is parking available?", vector_store=FakeStore([(document, 0.5)]), - model=FakeModel("\n INSUFFICIENT_EVIDENCE \n"), + model=model, ) self.assertEqual(result.answer, NO_MATCH_MESSAGE) self.assertEqual(result.sources, ()) self.assertEqual(result.retrieved_source_ids, ("review-1",)) self.assertTrue(result.abstained) + self.assertEqual(result.failure_reason, "clean_abstention") + self.assertFalse(result.repair_attempted) + self.assertEqual(len(model.prompts), 1) def test_accepts_cited_answer_with_standalone_insufficient_token_line( self, @@ -211,18 +306,46 @@ def test_rejects_insufficient_evidence_token_embedded_in_prose(self) -> None: id="review-1", ) + model = SequenceModel( + "The raw marker INSUFFICIENT_EVIDENCE must not be shown [1].", + "This response must never be requested [1].", + ) result = answer_question( "What do guests say about the crust?", vector_store=FakeStore([(document, 0.5)]), - model=FakeModel( - "The raw marker INSUFFICIENT_EVIDENCE must not be shown [1]." - ), + model=model, ) self.assertEqual(result.answer, CITATION_VALIDATION_MESSAGE) self.assertNotIn("INSUFFICIENT_EVIDENCE", result.answer) self.assertEqual(result.sources, ()) self.assertFalse(result.abstained) + self.assertEqual(result.failure_reason, "embedded_control_token") + self.assertFalse(result.repair_attempted) + self.assertEqual(len(model.prompts), 1) + + def test_rejects_empty_remainder_without_attempting_repair(self) -> None: + document = Document( + page_content="The crust was perfectly crispy.", + metadata={"source_id": "review-1"}, + id="review-1", + ) + model = SequenceModel( + "INSUFFICIENT_EVIDENCE\nINSUFFICIENT_EVIDENCE", + "This response must never be requested [1].", + ) + + result = answer_question( + "What do guests say about the crust?", + vector_store=FakeStore([(document, 0.5)]), + model=model, + ) + + self.assertEqual(result.answer, CITATION_VALIDATION_MESSAGE) + self.assertEqual(result.sources, ()) + self.assertEqual(result.failure_reason, "invalid_remainder") + self.assertFalse(result.repair_attempted) + self.assertEqual(len(model.prompts), 1) def test_rejects_uncited_answer_after_removing_control_token_line(self) -> None: document = Document( @@ -240,6 +363,7 @@ def test_rejects_uncited_answer_after_removing_control_token_line(self) -> None: self.assertEqual(result.answer, CITATION_VALIDATION_MESSAGE) self.assertEqual(result.sources, ()) self.assertFalse(result.abstained) + self.assertEqual(result.failure_reason, "missing_citations") def test_does_not_call_model_when_filters_match_no_reviews(self) -> None: model = FakeModel() @@ -255,6 +379,7 @@ def test_does_not_call_model_when_filters_match_no_reviews(self) -> None: self.assertEqual(result.retrieved_source_ids, ()) self.assertFalse(result.abstained) self.assertEqual(model.prompts, []) + self.assertEqual(result.failure_reason, "empty_retrieval") if __name__ == "__main__": diff --git a/tests/test_benchmark.py b/tests/test_benchmark.py index 42e17a6..09681be 100644 --- a/tests/test_benchmark.py +++ b/tests/test_benchmark.py @@ -191,7 +191,7 @@ def test_report_is_machine_readable_and_writes_matching_markdown(self) -> None: generated_at="2026-07-31T00:00:00Z", ) - self.assertEqual(report["schema_version"], 2) + self.assertEqual(report["schema_version"], 3) self.assertEqual(report["evaluation_set"]["case_count"], 2) self.assertEqual(report["evaluation_set"]["abstention_case_count"], 1) self.assertEqual(report["results"]["bm25_baseline"]["recall_at_k"], 0.0) @@ -213,6 +213,91 @@ def test_report_is_machine_readable_and_writes_matching_markdown(self) -> None: self.assertIn("Model-dependent results", markdown) self.assertIn("2026-07-31T00:00:00Z", markdown) + def test_report_serializes_raw_response_and_failure_diagnostics(self) -> None: + case = EvaluationCase( + case_id="answer", + question="Is it crisp?", + relevant_titles=("Crisp",), + reference_facts=(), + category="quality", + ) + observation = EvaluationObservation( + case_id="answer", + relevant_source_ids=frozenset({"a"}), + retrieved_source_ids=("a",), + cited_source_ids=(), + cited_text="", + answer="I could not produce an answer with valid citations.", + abstained=False, + outcome="citation_validation_rejection", + raw_model_response="The crust is crisp.", + repair_model_response="The crust is crisp [9].", + initial_failure_reason="missing_citations", + failure_reason="out_of_range_citation", + repair_attempted=True, + ) + + report = build_evaluation_report( + cases=(case,), + rag_metrics={ + "retrieval_recall": 1.0, + "citation_validity": 0.0, + "reference_term_support_proxy": 0.0, + "expected_action_accuracy": 0.0, + "answer_success_rate": 0.0, + "abstention_recall": 0.0, + "case_count": 1, + }, + semantic_metrics={ + "recall_at_k": 1.0, + "hit_rate_at_k": 1.0, + "mrr_at_k": 1.0, + "evaluated_case_count": 1, + "limit": 5, + }, + baseline_metrics={ + "recall_at_k": 0.0, + "hit_rate_at_k": 0.0, + "mrr_at_k": 0.0, + "evaluated_case_count": 1, + "limit": 5, + }, + observations=(observation,), + configuration={}, + provenance={}, + generated_at="2026-07-31T00:00:00Z", + ) + + serialized = report["observations"]["rag"][0] + self.assertEqual(serialized["raw_model_response"], "The crust is crisp.") + self.assertEqual(serialized["repair_model_response"], "The crust is crisp [9].") + self.assertEqual(serialized["initial_failure_reason"], "missing_citations") + self.assertEqual(serialized["failure_reason"], "out_of_range_citation") + self.assertTrue(serialized["repair_attempted"]) + self.assertEqual( + report["diagnostics"]["initial_failure_reasons"], + {"missing_citations": 1}, + ) + self.assertEqual( + report["diagnostics"]["final_failure_reasons"], + {"out_of_range_citation": 1}, + ) + self.assertEqual(report["diagnostics"]["repair_attempt_count"], 1) + self.assertEqual(report["diagnostics"]["repair_success_count"], 0) + + with tempfile.TemporaryDirectory() as directory: + markdown_path = Path(directory) / "README.md" + write_evaluation_report( + report, + json_path=Path(directory) / "report.json", + markdown_path=markdown_path, + ) + markdown = markdown_path.read_text(encoding="utf-8") + + self.assertIn("Answer diagnostics", markdown) + self.assertIn("missing_citations", markdown) + self.assertIn("out_of_range_citation", markdown) + class EvaluationSetQualityTest(unittest.TestCase): def test_bundled_set_has_broad_unique_coverage(self) -> None: diff --git a/tests/test_evaluation.py b/tests/test_evaluation.py index 5ac801d..9d357bb 100644 --- a/tests/test_evaluation.py +++ b/tests/test_evaluation.py @@ -28,6 +28,16 @@ def invoke(self, prompt: str) -> FakeResponse: return FakeResponse(self.response) +class EvaluationSequenceModel: + def __init__(self, *responses: str) -> None: + self.responses = list(responses) + self.prompts: list[str] = [] + + def invoke(self, prompt: str) -> FakeResponse: + self.prompts.append(prompt) + return FakeResponse(self.responses.pop(0)) + + class EvaluationStore: def __init__(self, *, return_matches: bool = True) -> None: self.search_calls = 0 @@ -156,6 +166,87 @@ def test_empty_retrieval_is_not_counted_as_model_abstention(self) -> None: self.assertEqual(model.prompts, []) self.assertEqual(metrics.abstention_accuracy, 0.0) + def test_pipeline_preserves_raw_rejection_and_repair_diagnostics(self) -> None: + case = EvaluationCase( + case_id="answer", + question="Is the crust crisp?", + relevant_titles=("Best pizza",), + reference_facts=( + ReferenceFact( + answer_terms=("crispy",), + source_terms=("perfectly crispy",), + ), + ), + ) + model = EvaluationSequenceModel( + "The crust is crispy.", + "The crust is crispy [1].", + ) + + metrics, observations = run_rag_evaluation( + (case,), vector_store=EvaluationStore(), model=model + ) + + observation = observations[0] + self.assertEqual(observation.outcome, "answered_after_repair") + self.assertEqual(observation.raw_model_response, "The crust is crispy.") + self.assertEqual(observation.repair_model_response, "The crust is crispy [1].") + self.assertEqual(observation.initial_failure_reason, "missing_citations") + self.assertIsNone(observation.failure_reason) + self.assertTrue(observation.repair_attempted) + self.assertEqual(metrics.answer_success_rate, 1.0) + + def test_pipeline_labels_an_abstention_returned_by_repair(self) -> None: + case = EvaluationCase( + case_id="abstain", + question="Is parking available?", + relevant_titles=(), + reference_facts=(), + should_abstain=True, + ) + model = EvaluationSequenceModel( + "Parking is available.", + "INSUFFICIENT_EVIDENCE", + ) + + metrics, observations = run_rag_evaluation( + (case,), vector_store=EvaluationStore(), model=model + ) + + self.assertEqual(observations[0].outcome, "model_abstention_after_repair") + self.assertTrue(observations[0].repair_attempted) + self.assertEqual(observations[0].failure_reason, "clean_abstention") + self.assertEqual(metrics.abstention_recall, 1.0) + + def test_pipeline_labels_a_rejection_after_failed_repair(self) -> None: + case = EvaluationCase( + case_id="answer", + question="Is the crust crisp?", + relevant_titles=("Best pizza",), + reference_facts=( + ReferenceFact( + answer_terms=("crispy",), + source_terms=("perfectly crispy",), + ), + ), + ) + model = EvaluationSequenceModel( + "The crust is crispy.", + "Still no citation.", + ) + + metrics, observations = run_rag_evaluation( + (case,), vector_store=EvaluationStore(), model=model + ) + + self.assertEqual( + observations[0].outcome, + "citation_validation_rejection_after_repair", + ) + self.assertTrue(observations[0].repair_attempted) + self.assertEqual(observations[0].failure_reason, "missing_citations") + self.assertEqual(metrics.answer_success_rate, 0.0) + def test_penalizes_missing_retrieval_invalid_citation_and_false_answer( self, ) -> None: