diff --git a/evaluations/evaluations/config.py b/evaluations/evaluations/config.py index 96c56232..29b96145 100644 --- a/evaluations/evaluations/config.py +++ b/evaluations/evaluations/config.py @@ -38,6 +38,18 @@ class ConversationInput(BaseModel): return "\n".join(f"{turn.speaker}: {turn.text}" for turn in self.turns) +class ScopedQuestion(BaseModel): + """A question and the databases it may draw on. + + The task function receives only a case's inputs, never its metadata, so a + per-case scope has to travel in the inputs. `sources=[]` covers no database + and `None` covers every one the client covers. + """ + + question: str + sources: list[str] | None = None + + @dataclass class DocumentPayload: uri: str diff --git a/evaluations/evaluations/datasets/__init__.py b/evaluations/evaluations/datasets/__init__.py index 4a968b5f..366b11f7 100644 --- a/evaluations/evaluations/datasets/__init__.py +++ b/evaluations/evaluations/datasets/__init__.py @@ -2,6 +2,7 @@ from evaluations.config import DatasetSpec from .frames import FRAMES_SPEC from .hotpotqa import HOTPOTQA_SPEC +from .multidb import MULTIDB_SPEC, MULTIDB_SURFACES_SPEC from .mtrag import ( MTRAG_CLAPNQ_LIVE_SPEC, MTRAG_CLAPNQ_LIVE_UNCOMPACTED_SPEC, @@ -24,6 +25,8 @@ DATASETS: dict[str, DatasetSpec] = { MTRAG_CLAPNQ_REWRITE_SPEC, MTRAG_CLAPNQ_LIVE_SPEC, MTRAG_CLAPNQ_LIVE_UNCOMPACTED_SPEC, + MULTIDB_SPEC, + MULTIDB_SURFACES_SPEC, ORB_TEXT_SPEC, ORB_MULTIMODAL_SPEC, ORB_MULTIMODAL_NEMOTRON_SPEC, diff --git a/evaluations/evaluations/datasets/multidb.py b/evaluations/evaluations/datasets/multidb.py index 6f6c62d9..02da0818 100644 --- a/evaluations/evaluations/datasets/multidb.py +++ b/evaluations/evaluations/datasets/multidb.py @@ -1,8 +1,14 @@ import hashlib -from collections.abc import Iterator +from collections.abc import Iterator, Mapping from dataclasses import dataclass from pathlib import Path +from typing import Any +from datasets import Dataset +from pydantic_evals import Case + +from evaluations.config import DatasetSpec, ScopedQuestion +from evaluations.evaluators.multidb import MultiDBScores from haiku.rag.client import HaikuRAG from haiku.rag.config.models import AppConfig @@ -416,3 +422,372 @@ if __name__ == "__main__": # pragma: no cover - operator entry point import asyncio asyncio.run(main()) + + +# --- question families ------------------------------------------------------- +# +# B-family runs against both capabilities, S-family against analysis only, so +# they are registered as two dataset keys: a surface question cannot pass under +# the RAG target, and scoring it there would be noise rather than a finding. + +UNIQUE_STATIONS = tuple( + s + for s in STATIONS + if s.name != SHARED_STATION and not any(s.name in pair for pair in NEAR_NAME_PAIRS) +) + + +def _row( + question_id: str, + family: str, + question: str, + *, + sources: list[str] | None = None, + scope: list[str] | None = None, + expected_value: float | None = None, + expected_values: list[float] | None = None, + distractor_value: float | None = None, + expected_ordered: list[str] | None = None, + forbidden_text: list[str] | None = None, + expected_sources: list[str] | None = None, + answerability: str = "ANSWERABLE", +) -> dict: + return { + "id": question_id, + "family": family, + "question": question, + "sources": sources, + "scope": scope, + "expected_value": expected_value, + "expected_values": expected_values, + "distractor_value": distractor_value, + "expected_ordered": expected_ordered, + "forbidden_text": forbidden_text, + "expected_sources": expected_sources, + "answerability": answerability, + } + + +def _twin_of(name: str) -> str: + for north, south in NEAR_NAME_PAIRS: + if name == north: + return south + if name == south: + return north + raise KeyError(name) + + +def behaviour_rows() -> list[dict]: + rows: list[dict] = [] + + # B1 — a fact held by exactly one database. + for i, s in enumerate(UNIQUE_STATIONS[:3]): + rows.append( + _row( + f"b1-{i}", + "B1", + f"At what elevation in metres does Station {s.name} sit?", + expected_value=s.elevation_m, + expected_sources=[s.database], + ) + ) + + # B2 — a fact from each database in one answer. RRF ties resolve to the + # order the databases are listed, so the order is rotated across instances: + # otherwise the family measures ordering rather than fusion. + order = list(DATABASE_NAMES) + for i, (north_name, south_name) in enumerate(NEAR_NAME_PAIRS): + north, south = station(north_name, NORTHERN), station(south_name, SOUTHERN) + higher = north if north.elevation_m > south.elevation_m else south + rotated = order[i % len(order) :] + order[: i % len(order)] + rows.append( + _row( + f"b2-{i}", + "B2", + f"Which sits higher, Station {north.name} or Station {south.name}, " + "and at what elevation in metres?", + sources=rotated, + expected_value=higher.elevation_m, + expected_sources=[NORTHERN, SOUTHERN], + ) + ) + + # B3 — the near-name distractor. Gold present and the twin's number absent, + # because a hedge naming both is the failure this family exists to catch. + for i, (north_name, south_name) in enumerate(NEAR_NAME_PAIRS): + for name, db in ((north_name, NORTHERN), (south_name, SOUTHERN)): + s = station(name, db) + twin = station(_twin_of(name), SOUTHERN if db == NORTHERN else NORTHERN) + rows.append( + _row( + f"b3-elev-{i}-{db}", + "B3", + f"At what elevation in metres does Station {s.name} sit?", + expected_value=s.elevation_m, + distractor_value=twin.elevation_m, + expected_sources=[db], + ) + ) + north = station(north_name, NORTHERN) + twin = station(south_name, SOUTHERN) + rows.append( + _row( + f"b3-year-{i}", + "B3", + f"In what year was Station {north.name} commissioned?", + expected_value=north.commissioned, + distractor_value=twin.commissioned, + expected_sources=[NORTHERN], + ) + ) + + # B4 — scoped. Half the instances ask about a near-name pair member with the + # twin excluded, so honouring scope costs the model the other strong match. + for i, (north_name, south_name) in enumerate(NEAR_NAME_PAIRS): + for name, db in ((north_name, NORTHERN), (south_name, SOUTHERN)): + s = station(name, db) + rows.append( + _row( + f"b4-pair-{i}-{db}", + "B4", + f"At what elevation in metres does Station {s.name} sit?", + sources=[db], + scope=[db], + expected_value=s.elevation_m, + expected_sources=[db], + ) + ) + for i, s in enumerate(UNIQUE_STATIONS[:4]): + rows.append( + _row( + f"b4-unique-{i}", + "B4", + f"In what year was Station {s.name} commissioned?", + sources=[s.database], + scope=[s.database], + expected_value=s.commissioned, + expected_sources=[s.database], + ) + ) + + # B5 — one entity in both databases, unscoped: both must surface and both + # must be attributed. + north, south = station(SHARED_STATION, NORTHERN), station(SHARED_STATION, SOUTHERN) + rows.append( + _row( + "b5-years", + "B5", + f"In what years was Station {SHARED_STATION} commissioned? There is a " + "station of that name in more than one programme.", + expected_values=[north.commissioned, south.commissioned], + expected_sources=[NORTHERN, SOUTHERN], + ) + ) + rows.append( + _row( + "b5-elevations", + "B5", + f"At what elevations do the stations named {SHARED_STATION} sit?", + expected_values=[north.elevation_m, south.elevation_m], + expected_sources=[NORTHERN, SOUTHERN], + ) + ) + rows.append( + _row( + "b5-programmes", + "B5", + f"Which programmes operate a Station {SHARED_STATION}?", + expected_ordered=[NORTHERN_PROGRAMME], + expected_sources=[NORTHERN, SOUTHERN], + ) + ) + + # B6 — absent stations. Judge-scored, because a deterministic refusal matcher + # keys on phrasing the model may never use. + for i, absent in enumerate(("Cormorant", "Razorbill", "Kittiwake")): + rows.append( + _row( + f"b6-{i}", + "B6", + f"At what elevation in metres does Station {absent} sit?", + answerability="UNANSWERABLE", + ) + ) + + # B7 — empty scope covers no database, so there is no evidence to answer + # from. API-only: the CLI has no --sources. + for i, s in enumerate(UNIQUE_STATIONS[:3]): + rows.append( + _row( + f"b7-{i}", + "B7", + f"At what elevation in metres does Station {s.name} sit?", + sources=[], + scope=[], + answerability="UNANSWERABLE", + ) + ) + return rows + + +def surface_rows() -> list[dict]: + rows: list[dict] = [] + + # S1 — inventory per database. Counts differ, so a count cannot be right by + # luck while attribution is wrong. + for name in DATABASE_NAMES: + rows.append( + _row( + f"s1-{name}", + "S1", + f"How many documents does the {name} database hold?", + expected_value=len(documents_for(name)), + ) + ) + + # S2 — a scoped search per programme, run in code. + for programme, db in ( + (NORTHERN_PROGRAMME, NORTHERN), + (SOUTHERN_PROGRAMME, SOUTHERN), + ): + for model in ("Vaisala WXT536", "Young 81000V"): + count = sum(1 for s in stations_in(db) if s.instrument == model) + rows.append( + _row( + f"s2-{db}-{model.split()[0].lower()}", + "S2", + f"How many stations in the {programme} use a {model}?", + expected_value=count, + ) + ) + + # S3 — a whole-document surface. No single chunk holds all twelve readings + # (asserted at build time), so the total cannot come from search alone. + for name, db in (("Kestrel", NORTHERN), ("Auk", NORTHERN), ("Snowcap", SOUTHERN)): + s = station(name, db) + rows.append( + _row( + f"s3-{db}-{s.slug}", + "S3", + f"What is the total of the twelve monthly mean wind speeds in the " + f"report for Station {s.name} in the {db} database?", + expected_value=s.readings_total, + ) + ) + + # S4 — structure: the table's row count, which needs the itemised document. + for name, db in (("Kestrel", NORTHERN), ("Albatross", SOUTHERN)): + s = station(name, db) + rows.append( + _row( + f"s4-{db}-{s.slug}", + "S4", + f"How many data rows does the monthly readings table in the report " + f"for Station {s.name} have?", + expected_value=len(MONTHS), + ) + ) + + # S5 — the outline, in document order. + for name, db in (("Kestrel", NORTHERN), ("Snowcap", SOUTHERN)): + s = station(name, db) + rows.append( + _row( + f"s5-{db}-{s.slug}", + "S5", + f"List the section headings of the report for Station {s.name}, in " + "the order they appear.", + expected_ordered=[ + "Overview", + "Instruments", + "Measurements", + "Maintenance", + ], + ) + ) + + # S6 — per-document metadata, with the twin's uri forbidden. + for north_name, south_name in NEAR_NAME_PAIRS[:2]: + s = station(north_name, NORTHERN) + twin = station(south_name, SOUTHERN) + rows.append( + _row( + f"s6-{s.slug}", + "S6", + f"What is the uri of the document titled {s.title!r}, and which " + "database holds it?", + expected_ordered=[s.uri], + forbidden_text=[twin.uri], + ) + ) + return rows + + +def load_behaviour_questions() -> Dataset: + return Dataset.from_list(behaviour_rows()) + + +def load_surface_questions() -> Dataset: + return Dataset.from_list(surface_rows()) + + +def build_multidb_case(index: int, row: Mapping[str, Any]) -> Case[Any, Any, dict]: + """One case. A scope travels in the inputs, since the task function receives + a case's inputs and never its metadata.""" + question = str(row["question"]) + sources = row["sources"] + inputs: str | ScopedQuestion = ( + question + if sources is None + else ScopedQuestion(question=question, sources=list(sources)) + ) + metadata = { + "question_id": str(row["id"]), + "family": str(row["family"]), + "case_index": str(index), + "answerability": row["answerability"], + } + for key in ( + "expected_value", + "expected_values", + "distractor_value", + "expected_ordered", + "forbidden_text", + "expected_sources", + "scope", + ): + if row[key] is not None: + metadata[key] = row[key] + return Case( + name=f"{index}_{row['id']}", + inputs=inputs, + expected_output=None, + metadata=metadata, + ) + + +def _unused_document_loader() -> Dataset: # pragma: no cover - never called + raise RuntimeError( + "the multi-database corpus is built by build_databases(); run with --skip-db" + ) + + +MULTIDB_SPEC = DatasetSpec( + key="multidb", + db_filename="multidb_northern.lancedb", + document_loader=_unused_document_loader, + document_mapper=lambda _row: None, + qa_loader=load_behaviour_questions, + qa_case_builder=build_multidb_case, + qa_evaluator=MultiDBScores(), +) + +MULTIDB_SURFACES_SPEC = DatasetSpec( + key="multidb_surfaces", + db_filename="multidb_northern.lancedb", + document_loader=_unused_document_loader, + document_mapper=lambda _row: None, + qa_loader=load_surface_questions, + qa_case_builder=build_multidb_case, + qa_evaluator=MultiDBScores(), +) diff --git a/evaluations/evaluations/evaluators/multidb.py b/evaluations/evaluations/evaluators/multidb.py new file mode 100644 index 00000000..97c0691e --- /dev/null +++ b/evaluations/evaluations/evaluators/multidb.py @@ -0,0 +1,156 @@ +import re +from collections.abc import Sequence +from dataclasses import dataclass + +from pydantic_evals.evaluators import Evaluator, EvaluatorContext + +_NUMBER = re.compile(r"-?\d[\d,]*(?:\.\d+)?") + + +def numbers_in(text: str) -> set[float]: + """Every number in the text, comma separators removed. + + Answers are scored by extraction rather than string matching: "1,240 metres", + "1240 m" and a sentence around either are all legitimate. + """ + found: set[float] = set() + for match in _NUMBER.finditer(text or ""): + try: + found.add(float(match.group().replace(",", ""))) + except ValueError: # pragma: no cover - the pattern only matches numbers + continue + return found + + +def _as_floats( + single: float | int | None, many: Sequence[float | int] | None +) -> set[float]: + values: list[float | int] = [] if single is None else [single] + values.extend(many or ()) + return {float(v) for v in values} + + +def cited_sources(ctx: EvaluatorContext) -> list[str]: + return [s for s in (ctx.attributes.get("cited_sources") or []) if s] + + +@dataclass +class NumericAnswer(Evaluator): + """The gold number is present and the distractor's is absent. + + Presence alone is the wrong assertion: "Station Kestrel sits at either 1240 m + or 2310 m" contains the gold value while demonstrating exactly the confusion + the near-name pair exists to provoke. Reads `expected_value` and optional + `distractor_value` from case metadata. + """ + + def get_default_evaluation_name(self) -> str: + return "answer_correct" + + def evaluate(self, ctx: EvaluatorContext) -> dict[str, bool]: + meta = ctx.metadata or {} + expected = _as_floats(meta.get("expected_value"), meta.get("expected_values")) + if not expected: + return {} + found = numbers_in(str(ctx.output)) + forbidden = _as_floats( + meta.get("distractor_value"), meta.get("distractor_values") + ) + correct = expected <= found and not (forbidden & found) + return {"answer_correct": correct} + + +@dataclass +class AttributionGate(Evaluator): + """The databases cited are exactly the databases that hold the answer. + + A hard gate: `cited_map` scores URIs and would pass an answer attributed to + the wrong database, which is the failure mode this dataset exists for. + """ + + def get_default_evaluation_name(self) -> str: + return "attribution_correct" + + def evaluate(self, ctx: EvaluatorContext) -> dict[str, bool]: + meta = ctx.metadata or {} + expected = meta.get("expected_sources") + if expected is None: + return {} + return {"attribution_correct": set(cited_sources(ctx)) == set(expected)} + + +@dataclass +class ScopeGate(Evaluator): + """Nothing outside a scoped question's `sources` is cited. + + A hard gate, and separate from attribution: a case can cite the right + database and still have reached outside its scope to get there. + """ + + def get_default_evaluation_name(self) -> str: + return "scope_honoured" + + def evaluate(self, ctx: EvaluatorContext) -> dict[str, bool]: + meta = ctx.metadata or {} + scope = meta.get("scope") + if scope is None: + return {} + return {"scope_honoured": set(cited_sources(ctx)) <= set(scope)} + + +@dataclass +class TextAnswer(Evaluator): + """Required strings appear, in order, and forbidden strings do not. + + Order matters for the headings case, where the outline is only right if the + sections come back in document order. + """ + + def get_default_evaluation_name(self) -> str: + return "answer_correct" + + def evaluate(self, ctx: EvaluatorContext) -> dict[str, bool]: + meta = ctx.metadata or {} + required = meta.get("expected_ordered") + if required is None: + return {} + haystack = str(ctx.output).lower() + cursor = 0 + for needle in required: + found = haystack.find(str(needle).lower(), cursor) + if found < 0: + return {"answer_correct": False} + cursor = found + len(str(needle)) + for forbidden in meta.get("forbidden_text") or []: + if str(forbidden).lower() in haystack: + return {"answer_correct": False} + return {"answer_correct": True} + + +@dataclass +class MultiDBScores(Evaluator): + """Every deterministic score for a case, in one evaluator. + + `DatasetSpec.qa_evaluator` takes a single evaluator and replaces the judge + when set, so the scorers are composed here rather than listed. Each abstains + on cases whose metadata does not ask for it. + """ + + def evaluate(self, ctx: EvaluatorContext) -> dict[str, bool]: + meta = ctx.metadata or {} + numeric = bool(meta.get("expected_value") or meta.get("expected_values")) + textual = meta.get("expected_ordered") is not None + if numeric and textual: + raise ValueError( + f"case {meta.get('question_id')!r} asks for both a numeric and an " + "ordered-text answer; they share the answer_correct key" + ) + scores: dict[str, bool] = {} + for evaluator in ( + NumericAnswer(), + TextAnswer(), + AttributionGate(), + ScopeGate(), + ): + scores.update(evaluator.evaluate(ctx)) + return scores diff --git a/evaluations/evaluations/qa.py b/evaluations/evaluations/qa.py index fdf8d0dd..d6a5e271 100644 --- a/evaluations/evaluations/qa.py +++ b/evaluations/evaluations/qa.py @@ -15,7 +15,7 @@ from evaluations.capability_runner import ( run_capability_conversation, run_capability_question, ) -from evaluations.config import ConversationInput, DatasetSpec +from evaluations.config import ConversationInput, DatasetSpec, ScopedQuestion from evaluations.evaluators import ( ANSWER_EQUIVALENCE_RUBRIC, REFUSAL_ELIGIBLE_LABELS, @@ -344,10 +344,15 @@ async def run_qa_benchmark( name=spec.key, cases=cases, evaluators=evaluators ) - async def answer_question(inputs: str | ConversationInput) -> str: + async def answer_question(inputs: str | ConversationInput | ScopedQuestion) -> str: + sources: list[str] | None = None if isinstance(inputs, ConversationInput): question = inputs.question message_history = prefix_to_messages(inputs.prefix) + elif isinstance(inputs, ScopedQuestion): + question = inputs.question + sources = inputs.sources + message_history = None else: question = inputs message_history = None @@ -359,6 +364,7 @@ async def run_qa_benchmark( capability_model=run.capability_model, document_filter=document_filter, message_history=message_history, + sources=sources, ) set_eval_attribute("cited_uris", result.cited_uris) set_eval_attribute("cited_chunk_ids", result.cited_chunk_ids) diff --git a/evaluations/tests/test_multidb_cases.py b/evaluations/tests/test_multidb_cases.py new file mode 100644 index 00000000..3420e63d --- /dev/null +++ b/evaluations/tests/test_multidb_cases.py @@ -0,0 +1,113 @@ +from collections import Counter + +from evaluations.config import ScopedQuestion +from evaluations.datasets.multidb import ( + DATABASE_NAMES, + NEAR_NAME_PAIRS, + behaviour_rows, + build_multidb_case, + surface_rows, +) + + +def by_family(rows): + grouped = {} + for row in rows: + grouped.setdefault(row["family"], []).append(row) + return grouped + + +def test_gate_families_carry_enough_instances(): + """B3 and B4 are pass/fail gates, so a handful of cases is not evidence of + absence. Everything else is a rate and three is enough.""" + families = by_family(behaviour_rows()) + assert len(families["B3"]) >= 8 + assert len(families["B4"]) >= 8 + + +def test_half_the_scoped_cases_exclude_the_better_match(): + """If the scoped-in database always holds the most relevant content, the + model answers correctly without honouring scope and no-leakage is trivially + satisfied. Half the B4 instances ask about a near-name pair member, whose + twin is the excluded strong match.""" + b4 = by_family(behaviour_rows())["B4"] + pair_cases = [r for r in b4 if r["id"].startswith("b4-pair-")] + assert len(pair_cases) >= len(b4) / 2 + twins = {name for pair in NEAR_NAME_PAIRS for name in pair} + for row in pair_cases: + assert any(f"Station {name}" in row["question"] for name in twins) + + +def test_single_source_families_expect_exactly_one_database(): + for row in behaviour_rows(): + if row["family"] in {"B1", "B3", "B4"}: + assert len(row["expected_sources"]) == 1 + + +def test_shared_entity_family_expects_both_databases(): + for row in by_family(behaviour_rows())["B5"]: + assert set(row["expected_sources"]) == {"northern", "southern"} + + +def test_cross_database_family_rotates_database_order(): + """RRF ties resolve to the order the databases are listed, so a fixed order + would make this family measure ordering rather than fusion.""" + orders = [tuple(row["sources"]) for row in by_family(behaviour_rows())["B2"]] + assert len(set(orders)) == len(orders) + for order in orders: + assert set(order) == set(DATABASE_NAMES) + + +def test_refusal_families_are_labelled_unanswerable(): + """The label is what makes RefusalJudge score them and what feeds refusal + precision and recall.""" + families = by_family(behaviour_rows()) + for family in ("B6", "B7"): + assert all(r["answerability"] == "UNANSWERABLE" for r in families[family]) + answerable = [ + r for f, rows in families.items() if f not in {"B6", "B7"} for r in rows + ] + assert all(r["answerability"] == "ANSWERABLE" for r in answerable) + + +def test_empty_scope_reaches_the_case_as_an_empty_list(): + """`sources=[]` covers nothing and must not collapse into None, which covers + everything.""" + rows = [r for r in behaviour_rows() if r["family"] == "B7"] + assert rows + for index, row in enumerate(rows): + case = build_multidb_case(index, row) + assert isinstance(case.inputs, ScopedQuestion) + assert case.inputs.sources == [] + assert case.metadata is not None + assert case.metadata["scope"] == [] + + +def test_unscoped_cases_pass_a_bare_question(): + row = next(r for r in behaviour_rows() if r["family"] == "B1") + case = build_multidb_case(0, row) + assert isinstance(case.inputs, str) + + +def test_no_case_asks_for_both_a_number_and_ordered_text(): + """The two scorers share the answer_correct key, and the composite raises + rather than letting one silently win.""" + for row in (*behaviour_rows(), *surface_rows()): + numeric = ( + row["expected_value"] is not None or row["expected_values"] is not None + ) + assert not (numeric and row["expected_ordered"] is not None) + + +def test_case_ids_are_unique_and_stable(): + rows = (*behaviour_rows(), *surface_rows()) + ids = [r["id"] for r in rows] + assert len(set(ids)) == len(ids) + assert Counter(r["family"] for r in surface_rows()).keys() >= { + "S1", + "S2", + "S3", + "S4", + "S5", + "S6", + } diff --git a/evaluations/tests/test_multidb_evaluators.py b/evaluations/tests/test_multidb_evaluators.py new file mode 100644 index 00000000..cf62a149 --- /dev/null +++ b/evaluations/tests/test_multidb_evaluators.py @@ -0,0 +1,119 @@ +from types import SimpleNamespace + +from evaluations.evaluators.multidb import ( + AttributionGate, + NumericAnswer, + ScopeGate, + TextAnswer, + numbers_in, +) + + +def ctx(output="", metadata=None, attributes=None): + return SimpleNamespace( + output=output, metadata=metadata or {}, attributes=attributes or {} + ) + + +def test_numbers_survive_comma_separators_and_units(): + assert numbers_in("Station Kestrel sits at 1,240 metres") == {1240.0} + assert numbers_in("1240 m") == {1240.0} + assert numbers_in("no digits here") == set() + + +def test_gold_number_counts_as_correct(): + result = NumericAnswer().evaluate( + ctx("It sits at 1240 metres.", {"expected_value": 1240}) + ) + assert result == {"answer_correct": True} + + +def test_hedging_between_gold_and_distractor_fails(): + """The failure the near-name pair exists to catch: a presence check would + pass this, since the gold value is in the answer.""" + result = NumericAnswer().evaluate( + ctx( + "Station Kestrel sits at either 1240 m or 2310 m.", + {"expected_value": 1240, "distractor_value": 2310}, + ) + ) + assert result == {"answer_correct": False} + + +def test_numeric_answer_abstains_without_a_gold_value(): + assert NumericAnswer().evaluate(ctx("anything", {})) == {} + + +def test_attribution_requires_the_exact_database_set(): + good = AttributionGate().evaluate( + ctx( + metadata={"expected_sources": ["northern"]}, + attributes={"cited_sources": ["northern"]}, + ) + ) + assert good == {"attribution_correct": True} + wrong = AttributionGate().evaluate( + ctx( + metadata={"expected_sources": ["northern"]}, + attributes={"cited_sources": ["southern"]}, + ) + ) + assert wrong == {"attribution_correct": False} + extra = AttributionGate().evaluate( + ctx( + metadata={"expected_sources": ["northern"]}, + attributes={"cited_sources": ["northern", "equipment"]}, + ) + ) + assert extra == {"attribution_correct": False} + + +def test_scope_allows_a_subset_and_rejects_an_outsider(): + """Citing fewer databases than allowed honours scope; citing one outside it + does not.""" + inside = ScopeGate().evaluate( + ctx( + metadata={"scope": ["northern", "southern"]}, + attributes={"cited_sources": ["northern"]}, + ) + ) + assert inside == {"scope_honoured": True} + outside = ScopeGate().evaluate( + ctx( + metadata={"scope": ["northern"]}, + attributes={"cited_sources": ["northern", "southern"]}, + ) + ) + assert outside == {"scope_honoured": False} + + +def test_empty_scope_is_honoured_only_by_citing_nothing(): + assert ScopeGate().evaluate( + ctx(metadata={"scope": []}, attributes={"cited_sources": []}) + ) == {"scope_honoured": True} + assert ScopeGate().evaluate( + ctx(metadata={"scope": []}, attributes={"cited_sources": ["northern"]}) + ) == {"scope_honoured": False} + + +def test_ordered_text_must_appear_in_order(): + meta = {"expected_ordered": ["Overview", "Instruments", "Measurements"]} + assert TextAnswer().evaluate( + ctx("Overview, then Instruments, then Measurements", meta) + ) == {"answer_correct": True} + assert TextAnswer().evaluate( + ctx("Measurements, then Overview, then Instruments", meta) + ) == {"answer_correct": False} + + +def test_forbidden_text_fails_even_when_required_text_is_present(): + result = TextAnswer().evaluate( + ctx( + "station://northern/kestrel and station://southern/kestrel-ridge", + { + "expected_ordered": ["station://northern/kestrel"], + "forbidden_text": ["station://southern/kestrel-ridge"], + }, + ) + ) + assert result == {"answer_correct": False}