From 45d7c5fde22add0ad0cbb4c133d888c4e7e1f742 Mon Sep 17 00:00:00 2001 From: Lia Date: Wed, 30 Sep 2026 10:01:34 +0000 Subject: [PATCH] =?UTF-8?q?=F0=9F=A7=AA=20test:=20Add=20Clean-Room=20Conte?= =?UTF-8?q?xt=20Selection=20Replay?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- evals/__init__.py | 1 + evals/context_selection/README.md | 211 +++++++++ evals/context_selection/__init__.py | 1 + evals/context_selection/cases.json | 615 +++++++++++++++++++++++++++ evals/context_selection/core.py | 191 +++++++++ evals/context_selection/providers.py | 187 ++++++++ evals/context_selection/replay.py | 342 +++++++++++++++ tests/test_context_selection.py | 479 +++++++++++++++++++++ 8 files changed, 2027 insertions(+) create mode 100644 evals/__init__.py create mode 100644 evals/context_selection/README.md create mode 100644 evals/context_selection/__init__.py create mode 100644 evals/context_selection/cases.json create mode 100644 evals/context_selection/core.py create mode 100644 evals/context_selection/providers.py create mode 100644 evals/context_selection/replay.py create mode 100644 tests/test_context_selection.py diff --git a/evals/__init__.py b/evals/__init__.py new file mode 100644 index 00000000..132c3332 --- /dev/null +++ b/evals/__init__.py @@ -0,0 +1 @@ +"""Isolated experiments, not production service modules.""" diff --git a/evals/context_selection/README.md b/evals/context_selection/README.md new file mode 100644 index 00000000..2ab8ef0d --- /dev/null +++ b/evals/context_selection/README.md @@ -0,0 +1,211 @@ +# Context selection experiment + +This is an independently authored, opt-in replay harness, not a production +search change. It tests whether selecting **answer-bearing evidence**, rather +than always filling a top-k window, can reduce context without losing facts, +exceptions, or opposing evidence. + +The motivation is GPT Researcher's context-filter write-up: +https://docs.gptr.dev/docs/gpt-researcher/gptr/context-filter + +No upstream implementation or fixtures were copied. BM25 is the standard +formula, with a deliberately simple NFKC/case-folded word tokenizer that retains +negation and identifiers. We previously inspected upstream source, so this is +not a claim of a formal legal clean-room process. The provider protocol is +interoperability, not a dependency on GPT Researcher or its classes. + +## Boundary and ownership + +```text +fixed, already-authorized passage pool + query + -> score once (BM25 / embeddings / Jev) + -> selection ablations using those same scores + -> original passage IDs and sources + evidence-retention metrics +``` + +The harness imports neither `app.config` nor application storage. Production +routes, schemas, chunking, caches, authorization and ingestion are untouched. +There is no server endpoint, database probe, scrape, writer LLM or automatic OCR. +The operator owns corpus authorization and permission to send its content to a +provider. Never use private production documents without that authorization. + +A candidate is one complete existing chunk. This experiment does not re-chunk +or summarize it. Keep candidate generation constant across strategies; later +chunk-size experiments need separately fingerprinted corpora. A selector cannot +recover evidence omitted by the retriever. + +Selection owns stable tie ordering, thresholds and output budgets. The HTTP +adapter owns response validation, shared concurrency and cancellation. Evaluation +labels are used only after selection and are never included in provider payloads. + +## Strategies + +| Strategy | Behavior | +| --- | --- | +| `all` | Original pool order, subject to the character budget, not top-k | +| `input-topk` | Original retrieval order, top-k and character budget | +| `bm25-topk` | Lexical scoring without a relevance threshold | +| `bm25-filter` | Lexical scoring, relative cutoff against the best score | +| `embeddings-topk` | Cosine ranking without a threshold | +| `embeddings-filter` | Cosine ranking with an absolute threshold | +| `jev-topk` | Usefulness ranking without a threshold | +| `jev-filter` | Usefulness ranking with an absolute threshold | + +`max_results` is a ceiling, not a requirement to fill slots. Relative BM25 +selection returns empty when no term matches. Successful Jev rejection returns +empty, never the first passages as a disguised fallback. A provider failure is +an **error**, not a successful lexical run attributed to the provider. Provider +errors make the command exit 1; skipped runs are explicitly counted. + +Thresholds are experimental settings, not demonstrated calibration. Jev scores +are expected rubric levels in `[0, 3]`, not answer-correctness probabilities. +BM25 and cosine scores are not on that scale. There is no evidence that one +threshold fits every query type, language or model. + +`max_chars` counts Python Unicode characters in complete selected passage +bodies. It is **not** a model token budget, UTF-8 byte budget or final prompt +budget. Citation wrappers also cost tokens. Oversized passages are skipped, +not truncated into apparently complete evidence. Optional exact-text +deduplication applies to scored strategies; `all` and `input-topk` remain +unmodified controls. It does not merge overlapping or paraphrased passages. + +## Run offline + +Python 3.12 is the repository's CI version. The local CLI uses only the standard +library. HTTP checks need `httpx` from `test_requirements.txt`. + +```sh +mkdir -p .venv +python3 -m evals.context_selection.replay \ + --max-results 2 --output .venv/context-offline.json +``` + +The default includes BM25 and explicitly **skips** embeddings and Jev. Even if +keys exist, no client is constructed without `--allow-network`. There is no +`.env` loading. Missing live configuration is also recorded as skipped, never +silently replaced with simulated model scores. + +The committed 10 cases and 82 passages are synthetic boundary fixtures: +refund exceptions, exact issue identifiers, negation, multi-facet answers, +conflicting dated sources, Unicode, duplicate passages, evidence after the +first 50 candidates, irrelevant pools, and empty pools. They establish behavior +and expose lexical blind spots. They are **not** a representative benchmark or +an independent estimate of answer quality. + +## Live replay, once credentials are provided + +Set these through the environment or secret tooling, not tracked files or PR +comments: + +- `CONTEXT_EMBEDDING_URL`: full OpenAI-compatible embeddings endpoint URL. +- `CONTEXT_EMBEDDING_MODEL`: explicit model name. +- `CONTEXT_EMBEDDING_KEY`: provider key. +- `TYPESAFE_API_KEY`: Jev provider key. +- `CONTEXT_JEV_MODEL`: defaults to `jev-latest`; pin a model when possible. +- `CONTEXT_JEV_URL`: optional, defaults to the System One endpoint. + +The URLs are operator-controlled endpoints, not untrusted user URL inputs. +TLS verification remains enabled and redirects are not followed. Secrets and +provider response bodies are not written to reports or printed on failures. +Queries and passages go to the configured provider only on explicitly enabled +runs. Score outputs include models, indices, call counts, reported input tokens +(or null if absent), and measured scoring time. Failed batches retain attempted +request counts and the subtotal of input tokens reported before failure, but +mark total usage unknown. No pricing is invented. + +```sh +python3 -m evals.context_selection.replay --allow-network \ + --max-results 2 --output .venv/context-live.json + +# Sweep selection settings without paying for the same scores again. +python3 -m evals.context_selection.replay \ + --reuse-scores .venv/context-live.json --jev-min 2.0 \ + --max-results 4 --output .venv/context-sweep.json +``` + +Corpus bytes and schema version must match for reuse. Scoring happens once per +case/provider; top-k and filtering reuse that result. Costs/calls must be read +from `score_runs`, not summed across selection rows. A saved-score replay +records zero new provider calls. Do not compare its inherited scoring latency +to a new network run. + +Jev uses a fixed number of worker tasks and a semaphore shared across calls on +that adapter. A batch deadline includes queue wait. Failure or cancellation +cancels and drains sibling tasks. We intentionally make no automatic retries; +rerun a failed experiment explicitly. The per-case candidate/character bounds +are enforced before any provider call. The default embedding batch contains +query plus candidate text; provider-specific batch limits may be lower and +should be reflected in the flags. Do not silently drop excess candidates. + +## Corpus and metrics + +`--corpus` accepts a JSON array of cases. Passage order is the original +retrieval order. IDs must be unique within a case; case IDs are globally unique. +Each required fact has an evidence label; any passage may support several. +Every required label must be present somewhere in the candidate pool, so +selection recall is not confused with retrieval recall. A `required: []` case +means no candidate contains useful answer evidence, not that the real-world +question has no answer. + +```json +[ + { + "id": "retention-example", + "query": "When can a record be deleted?", + "required": ["duration", "legal-hold"], + "passages": [ + {"id": "p1", "source": "policy:page1", "text": "Keep records for seven years.", "supports": ["duration"]}, + {"id": "p2", "source": "policy:page2", "text": "A legal hold prevents deletion.", "supports": ["legal-hold"]} + ] + } +] +``` + +Reports retain IDs and source handles, not query or passage bodies. They contain: + +- evidence precision: fraction of selected passages carrying annotated evidence; +- evidence recall: fraction of required labels retained; +- missing-evidence labels, including exceptions and contradicting sources; +- correct abstentions for pools with no annotated evidence; +- selected characters and passage counts; +- separate successful, skipped and error counts. + +Precision is null for empty selections, recall is null for no-evidence cases. +Macro summaries report their denominators and exclude null values, so inspect +missing-evidence and abstention counts alongside them. High precision alone can mean useful facts were lost. +This is label-based selection evaluation, not generated-answer faithfulness. + +## Verify this slice + +```sh +# The harness needs no application fixtures, database, or embedding initialization. +python -m pytest --noconftest tests/test_context_selection.py +python -m black --check evals/context_selection tests/test_context_selection.py +python -m compileall -q evals/context_selection tests/test_context_selection.py +``` + +The test file is also discovered by the repository's normal CI unit-test run. +MockTransport emulates external APIs; failure/cancellation tests exercise the +actual adapter tasks and cleanup, not a fake implementation of selection. + +## Acceptance before production integration + +1. Replay representative, authorized LibreChat web/file candidate pools with + human-reviewed evidence labels. Keep a held-out set for threshold selection. +2. Compare identical candidates, budgets, and writer models if answer generation + is added. Track identifier queries, broad questions, caveats, conflicts, + multilingual content and answer evidence late in the pool separately. +3. Require retained evidence and citation provenance, not token reduction alone. + Record provider latency, inference usage and downstream prompt-token counts + with the actual tokenizer before claiming speed or cost wins. +4. Exercise the real adapters with provider credentials. MockTransport tests + prove our expected HTTP contract, not current provider behavior or calibration. +5. Only then add an opt-in production selector. Existing search is the rollback + path. Auth/tenant/entity scope must be checked before inference egress, and + filtering must never substitute for whole-file content inspection. +6. Update the consuming agents contract to distinguish an honest empty evidence + set from service failure. Its current RAG reranker treats empty results as a + bad response and falls back; that must not be reused unchanged for filtering. + +No model quality, calibration, answer-quality, real-provider latency, downstream +cost reduction, or production no-regression claim follows from the offline run. diff --git a/evals/context_selection/__init__.py b/evals/context_selection/__init__.py new file mode 100644 index 00000000..5c454e14 --- /dev/null +++ b/evals/context_selection/__init__.py @@ -0,0 +1 @@ +"""Clean-room context-selection experiment.""" diff --git a/evals/context_selection/cases.json b/evals/context_selection/cases.json new file mode 100644 index 00000000..1238808a --- /dev/null +++ b/evals/context_selection/cases.json @@ -0,0 +1,615 @@ +[ + { + "id": "refund-exception", + "query": "When are refunds allowed, and what exception applies?", + "required": [ + "window", + "exception" + ], + "passages": [ + { + "id": "refund-exception:0", + "source": "overview", + "text": "Refund policy and refund requests are discussed in this refund overview.", + "supports": [] + }, + { + "id": "refund-exception:1", + "source": "terms", + "text": "Unused purchases are refundable for 30 days after payment.", + "supports": [ + "window" + ] + }, + { + "id": "refund-exception:2", + "source": "terms", + "text": "Except for activated licenses: after activation no refund is available.", + "supports": [ + "exception" + ] + }, + { + "id": "refund-exception:3", + "source": "marketing", + "text": "We prioritize customer happiness and a smooth purchasing experience.", + "supports": [] + } + ] + }, + { + "id": "identifier", + "query": "What fixes BUG_1842?", + "required": [ + "fix" + ], + "passages": [ + { + "id": "identifier:0", + "source": "index", + "text": "BUG_1842 BUG_1842 BUG_1842 issue tracker and release index.", + "supports": [] + }, + { + "id": "identifier:1", + "source": "release", + "text": "For BUG_1842, clear the stale lease before acquiring the replacement lock.", + "supports": [ + "fix" + ] + }, + { + "id": "identifier:2", + "source": "notes", + "text": "Fixes and bugs are recorded in our change history.", + "supports": [] + } + ] + }, + { + "id": "negation", + "query": "Does the retention rule require immediate deletion?", + "required": [ + "rule", + "hold" + ], + "passages": [ + { + "id": "negation:0", + "source": "overview", + "text": "Retention deletion: our guide discusses deletion and retention.", + "supports": [] + }, + { + "id": "negation:1", + "source": "policy", + "text": "No. Retention requires keeping the record for seven years.", + "supports": [ + "rule" + ] + }, + { + "id": "negation:2", + "source": "legal", + "text": "A legal hold prevents deletion even after the retention period expires.", + "supports": [ + "hold" + ] + } + ] + }, + { + "id": "multi-facet", + "query": "What caused the battery recall, and what are replacement tradeoffs?", + "required": [ + "cause", + "benefit", + "risk" + ], + "passages": [ + { + "id": "multi-facet:0", + "source": "news", + "text": "The battery recall has generated substantial discussion across the industry.", + "supports": [] + }, + { + "id": "multi-facet:1", + "source": "recall", + "text": "A separator defect can short-circuit cells and caused the recall.", + "supports": [ + "cause" + ] + }, + { + "id": "multi-facet:2", + "source": "replacement", + "text": "The replacement reduces fire risk by using a more stable chemistry.", + "supports": [ + "benefit" + ] + }, + { + "id": "multi-facet:3", + "source": "lab", + "text": "The safer replacement has 15 percent less capacity and adds 200 grams.", + "supports": [ + "risk" + ] + }, + { + "id": "multi-facet:4", + "source": "advert", + "text": "Read more battery recall news and battery recall information here.", + "supports": [] + } + ] + }, + { + "id": "conflicting-sources", + "query": "How long is the service retention period?", + "required": [ + "old", + "current" + ], + "passages": [ + { + "id": "conflicting-sources:0", + "source": "old-handbook", + "text": "The 2023 handbook states a retention period of 90 days.", + "supports": [ + "old" + ] + }, + { + "id": "conflicting-sources:1", + "source": "current-policy", + "text": "The 2026 policy supersedes the handbook: retention is now 30 days.", + "supports": [ + "current" + ] + }, + { + "id": "conflicting-sources:2", + "source": "menu", + "text": "Service retention period help center, documentation, pricing and support.", + "supports": [] + } + ] + }, + { + "id": "unicode", + "query": "保管 期限 は?", + "required": [ + "duration" + ], + "passages": [ + { + "id": "unicode:0", + "source": "navigation", + "text": "保管 期限 メニュー 情報 一覧", + "supports": [] + }, + { + "id": "unicode:1", + "source": "policy", + "text": "保管 期限 は 30日 です。", + "supports": [ + "duration" + ] + } + ] + }, + { + "id": "duplicates", + "query": "What is the upload limit?", + "required": [ + "size", + "exception" + ], + "passages": [ + { + "id": "duplicates:0", + "source": "mirror-a", + "text": "The upload limit is 15 megabytes.", + "supports": [ + "size" + ] + }, + { + "id": "duplicates:1", + "source": "mirror-b", + "text": "The upload limit is 15 megabytes.", + "supports": [ + "size" + ] + }, + { + "id": "duplicates:2", + "source": "policy", + "text": "Except for administrators, who may raise the upload limit.", + "supports": [ + "exception" + ] + }, + { + "id": "duplicates:3", + "source": "menu", + "text": "Upload limit settings and upload limit controls.", + "supports": [] + } + ] + }, + { + "id": "late-evidence", + "query": "What condition allows cancelling the contract?", + "required": [ + "condition" + ], + "passages": [ + { + "id": "late-evidence:0", + "source": "page-0", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:1", + "source": "page-1", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:2", + "source": "page-2", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:3", + "source": "page-3", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:4", + "source": "page-4", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:5", + "source": "page-5", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:6", + "source": "page-6", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:7", + "source": "page-7", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:8", + "source": "page-8", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:9", + "source": "page-9", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:10", + "source": "page-10", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:11", + "source": "page-11", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:12", + "source": "page-12", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:13", + "source": "page-13", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:14", + "source": "page-14", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:15", + "source": "page-15", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:16", + "source": "page-16", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:17", + "source": "page-17", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:18", + "source": "page-18", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:19", + "source": "page-19", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:20", + "source": "page-20", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:21", + "source": "page-21", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:22", + "source": "page-22", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:23", + "source": "page-23", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:24", + "source": "page-24", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:25", + "source": "page-25", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:26", + "source": "page-26", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:27", + "source": "page-27", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:28", + "source": "page-28", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:29", + "source": "page-29", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:30", + "source": "page-30", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:31", + "source": "page-31", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:32", + "source": "page-32", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:33", + "source": "page-33", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:34", + "source": "page-34", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:35", + "source": "page-35", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:36", + "source": "page-36", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:37", + "source": "page-37", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:38", + "source": "page-38", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:39", + "source": "page-39", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:40", + "source": "page-40", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:41", + "source": "page-41", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:42", + "source": "page-42", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:43", + "source": "page-43", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:44", + "source": "page-44", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:45", + "source": "page-45", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:46", + "source": "page-46", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:47", + "source": "page-47", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:48", + "source": "page-48", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:49", + "source": "page-49", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:50", + "source": "page-50", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:51", + "source": "page-51", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:52", + "source": "page-52", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:53", + "source": "page-53", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:54", + "source": "page-54", + "text": "Contract cancellation and contract conditions: introductory navigation.", + "supports": [] + }, + { + "id": "late-evidence:55", + "source": "terms-last", + "text": "Cancellation is allowed only after written notice and a 14-day cure period.", + "supports": [ + "condition" + ] + } + ] + }, + { + "id": "unanswerable", + "query": "What is the asteroid surface temperature?", + "required": [], + "passages": [ + { + "id": "unanswerable:0", + "source": "sports", + "text": "The local team won the regional championship.", + "supports": [] + }, + { + "id": "unanswerable:1", + "source": "baking", + "text": "Bread preparation uses flour, yeast, and warm water.", + "supports": [] + } + ] + }, + { + "id": "empty-pool", + "query": "What does the unavailable policy require?", + "required": [], + "passages": [] + } +] diff --git a/evals/context_selection/core.py b/evals/context_selection/core.py new file mode 100644 index 00000000..2dcad1fe --- /dev/null +++ b/evals/context_selection/core.py @@ -0,0 +1,191 @@ +"""Experimental selection contracts. No application imports or provider calls.""" + +from collections import Counter +from dataclasses import dataclass +import math +import re +import unicodedata + + +@dataclass(frozen=True) +class Passage: + id: str + source: str + text: str + supports: tuple[str, ...] = () + + +@dataclass(frozen=True) +class Case: + id: str + query: str + required: tuple[str, ...] + passages: tuple[Passage, ...] + + def __post_init__(self): + if not self.id or not self.query.strip(): + raise ValueError("case id and query must be nonempty") + ids = [p.id for p in self.passages] + if len(ids) != len(set(ids)) or any(not value for value in ids): + raise ValueError("passage ids must be nonempty and unique within a case") + if len(self.required) != len(set(self.required)): + raise ValueError("required evidence labels must be unique") + labels = set(self.required) + for passage in self.passages: + if not passage.source or not passage.text.strip(): + raise ValueError("passages require a source and nonempty text") + if not set(passage.supports) <= labels: + raise ValueError("passage references an unknown evidence label") + available = {label for p in self.passages for label in p.supports} + if available != labels: + raise ValueError( + "every required label needs evidence in the candidate pool" + ) + + +@dataclass(frozen=True) +class Budget: + max_results: int = 10 + max_chars: int = 12000 + + def __post_init__(self): + if self.max_results < 0 or self.max_chars < 0: + raise ValueError("budgets must be nonnegative") + + +def tokens(text: str) -> list[str]: + """Preserve identifiers, negations and Unicode; no English-only stemming.""" + return re.findall(r"\w+", unicodedata.normalize("NFKC", text).casefold()) + + +def bm25(query: str, passages: tuple[Passage, ...]) -> tuple[float, ...]: + """Standard BM25, k1=1.5 and b=0.75, with pool-local document frequency.""" + frequencies = [Counter(tokens(p.text)) for p in passages] + if not frequencies: + return () + lengths = [sum(frequency.values()) for frequency in frequencies] + average = sum(lengths) / len(lengths) or 1.0 + terms = set(tokens(query)) + document_frequency = Counter( + term for frequency in frequencies for term in terms if term in frequency + ) + scores = [] + for frequency, length in zip(frequencies, lengths): + score = 0.0 + for term in sorted(terms): + count = frequency[term] + if not count: + continue + df = document_frequency[term] + idf = math.log1p((len(passages) - df + 0.5) / (df + 0.5)) + score += ( + idf * count * 2.5 / (count + 1.5 * (0.25 + 0.75 * length / average)) + ) + scores.append(score) + return tuple(scores) + + +def number(value: object) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError("expected a finite number") + result = float(value) + if not math.isfinite(result): + raise ValueError("expected a finite number") + return result + + +def cosine(left: list[float], right: list[float]) -> float: + if not left or len(left) != len(right): + raise ValueError("embedding dimensions must match and be nonempty") + left = [number(value) for value in left] + right = [number(value) for value in right] + a = math.sqrt(math.fsum(value * value for value in left)) + b = math.sqrt(math.fsum(value * value for value in right)) + if not math.isfinite(a * b) or a == 0 or b == 0: + raise ValueError("embedding norms must be finite and nonzero") + return max(-1.0, min(1.0, math.fsum(x * y for x, y in zip(left, right)) / (a * b))) + + +def select( + passages: tuple[Passage, ...], + budget: Budget, + scores: tuple[float, ...] | None = None, + *, + minimum: float | None = None, + relative: float | None = None, + deduplicate: bool = False, +) -> tuple[Passage, ...]: + """Select whole original passages; never pad an honestly empty selection. + + Stable ties use input order. The character budget includes passage bodies, + not citation wrappers, and is not a tokenizer-derived prompt-token budget. + """ + if scores is not None: + if len(scores) != len(passages): + raise ValueError("one score is required per passage") + scores = tuple(number(score) for score in scores) + if minimum is not None: + minimum = number(minimum) + if relative is not None: + relative = number(relative) + if not 0 <= relative <= 1: + raise ValueError("relative threshold must be in [0, 1]") + if (minimum is not None or relative is not None) and scores is None: + raise ValueError("thresholds require scores") + order = ( + sorted(range(len(passages)), key=lambda i: (-scores[i], i)) + if scores is not None + else range(len(passages)) + ) + best = max(scores, default=0.0) if scores is not None else 0.0 + chosen = [] + used = 0 + seen = set() + for index in order: + if len(chosen) >= budget.max_results: + break + if scores is not None: + if minimum is not None and scores[index] < minimum: + continue + if relative is not None and (best <= 0 or scores[index] < relative * best): + continue + passage = passages[index] + # Exact duplicates only. Paraphrase/overlap merging needs separate evidence. + if deduplicate and passage.text in seen: + continue + if used + len(passage.text) > budget.max_chars: + continue + chosen.append(passage) + seen.add(passage.text) + used += len(passage.text) + return tuple(chosen) + + +def measure(case: Case, selected: tuple[Passage, ...]) -> dict: + originals = {p.id: p for p in case.passages} + if len({p.id for p in selected}) != len(selected) or any( + originals.get(p.id) != p for p in selected + ): + raise ValueError( + "selection must preserve original passage identity and content" + ) + covered = {label for p in selected for label in p.supports} + missing = sorted(set(case.required) - covered) + return { + "selected_ids": [p.id for p in selected], + "sources": [p.source for p in selected], + "selected_count": len(selected), + "input_chars": sum(len(p.text) for p in case.passages), + "selected_chars": sum(len(p.text) for p in selected), + "evidence_precision": ( + sum(bool(p.supports) for p in selected) / len(selected) + if selected + else None + ), + "evidence_recall": ( + len(covered) / len(case.required) if case.required else None + ), + "missing_evidence": missing, + "complete_evidence": not missing, + "correct_abstention": not selected if not case.required else None, + } diff --git a/evals/context_selection/providers.py b/evals/context_selection/providers.py new file mode 100644 index 00000000..4078b157 --- /dev/null +++ b/evals/context_selection/providers.py @@ -0,0 +1,187 @@ +"""Provider contracts for the experiment, never imported by the production API.""" + +import asyncio +from dataclasses import dataclass + +import httpx + +from .core import Passage, cosine, number + + +class ProviderError(Exception): + """Only this stable code is reported; provider bodies may contain secrets/text.""" + + def __init__(self, code: str): + self.code = code + super().__init__(code) + + +@dataclass(frozen=True) +class Scored: + values: tuple[float, ...] + model: str + requests: int + input_tokens: int | None + + +class HTTPScorer: + def __init__( + self, + client: httpx.AsyncClient, + endpoint: str, + model: str, + key: str, + *, + concurrency: int = 4, + timeout: float = 30.0, + ): + if concurrency < 1 or number(timeout) <= 0 or not model or not key: + raise ValueError("provider configuration is incomplete or invalid") + url = httpx.URL(endpoint) + if ( + url.scheme not in ("http", "https") + or not url.host + or url.username + or url.password + ): + raise ValueError("endpoint must be an HTTP(S) URL without credentials") + self.client = client + self.endpoint = endpoint + self.model = model + self.key = key + self.concurrency = concurrency + self.slots = asyncio.Semaphore(concurrency) + self.timeout = timeout + self.requests = 0 + self.reported_input_tokens = 0 + + async def _post(self, payload: dict, token_field: str) -> tuple[dict, int | None]: + async with self.slots: + self.requests += 1 + try: + response = await self.client.post( + self.endpoint, + headers={"Authorization": f"Bearer {self.key}"}, + json=payload, + timeout=self.timeout, + follow_redirects=False, + ) + except httpx.HTTPError: + raise ProviderError("transport_error") from None + if response.status_code != 200: + raise ProviderError(f"http_{response.status_code}") + try: + body = response.json() + if not isinstance(body, dict): + raise ValueError + usage = body.get("usage", {}) + if not isinstance(usage, dict): + raise ValueError + count = usage.get(token_field) + if count is not None and ( + isinstance(count, bool) or not isinstance(count, int) or count < 0 + ): + raise ValueError + if count is not None: + self.reported_input_tokens += count + return body, count + except (ValueError, TypeError): + raise ProviderError("invalid_response") from None + + async def embeddings(self, query: str, passages: tuple[Passage, ...]) -> Scored: + """One embedding batch, with explicit response-index validation.""" + if not passages: + return Scored((), self.model, 0, 0) + async with asyncio.timeout(self.timeout): + body, usage = await self._post( + {"model": self.model, "input": [query, *[p.text for p in passages]]}, + "prompt_tokens", + ) + try: + count = len(passages) + 1 + rows = body["data"] + if not isinstance(rows, list) or len(rows) != count: + raise ValueError + vectors = [None] * count + for row in rows: + index = row["index"] + if ( + isinstance(index, bool) + or not isinstance(index, int) + or not 0 <= index < count + or vectors[index] is not None + or not isinstance(row["embedding"], list) + ): + raise ValueError + vectors[index] = row["embedding"] + values = tuple(cosine(vectors[0], vector) for vector in vectors[1:]) + # An explicit model mismatch invalidates a supposedly fixed-space replay. + actual_model = body.get("model", self.model) + if actual_model != self.model: + raise ValueError + return Scored(values, self.model, 1, usage) + except (KeyError, TypeError, ValueError, OverflowError): + raise ProviderError("invalid_embeddings") from None + + async def jev(self, query: str, passages: tuple[Passage, ...]) -> Scored: + """Use fixed worker tasks, not one queued task per passage. + + Any failure cancels and drains siblings. Cancellation is not converted to + fallback. Successful threshold rejection remains an honestly empty result. + The shared semaphore also bounds overlapping score() calls on this adapter. + """ + values = [0.0] * len(passages) + usages = [None] * len(passages) + models = [self.model] * len(passages) + remaining = iter(enumerate(passages)) + + async def worker(): + for index, passage in remaining: + body, usages[index] = await self._post( + { + "model": self.model, + "state": passage.text, + "questions": { + "usefulness": { + "type": "score", + "instructions": ( + "Rate evidence useful for answering this question, " + "including necessary conditions, exceptions, or " + f"contradicting facts: {query}" + ), + "criteria": [ + "No information bearing on the question", + "Topic overlap without answer evidence", + "Evidence for part of the question or a necessary caveat", + "Specific evidence for the main requested facts", + ], + } + }, + }, + "input_tokens", + ) + try: + score = number(body["answers"]["usefulness"]["score"]) + models[index] = body.get("model", self.model) + if not 0 <= score <= 3 or not isinstance(models[index], str): + raise ValueError + values[index] = score + except (KeyError, TypeError, ValueError): + raise ProviderError("invalid_usefulness") from None + + try: + async with asyncio.timeout(self.timeout): + async with asyncio.TaskGroup() as group: + for _ in range(min(self.concurrency, len(passages))): + group.create_task(worker()) + except ExceptionGroup as errors: + # Suppress the exception group: nested messages are provider-controlled. + for error in errors.exceptions: + if isinstance(error, ProviderError): + raise error from None + raise ProviderError("scoring_failed") from None + if len(set(models)) > 1: + raise ProviderError("inconsistent_model") + total = sum(usages) if all(value is not None for value in usages) else None + model = models[0] if models else self.model + return Scored(tuple(values), model, len(passages), total) diff --git a/evals/context_selection/replay.py b/evals/context_selection/replay.py new file mode 100644 index 00000000..e28b1bb8 --- /dev/null +++ b/evals/context_selection/replay.py @@ -0,0 +1,342 @@ +"""Replay fixed, already-authorized passages. Run with python -m ...replay.""" + +import argparse +import asyncio +from dataclasses import asdict +import hashlib +import json +import os +import platform +from pathlib import Path +import time + +from .core import Budget, Case, Passage, bm25, measure, number, select + + +def load_cases(path: Path) -> tuple[tuple[Case, ...], str]: + raw = path.read_bytes() + try: + data = json.loads(raw) + cases = [] + for item in data: + if not isinstance(item["id"], str) or not isinstance(item["query"], str): + raise ValueError + required = item["required"] + if not isinstance(required, list) or any( + not isinstance(x, str) for x in required + ): + raise ValueError + passages = [] + for p in item["passages"]: + if any(not isinstance(p[key], str) for key in ("id", "source", "text")): + raise ValueError + supports = p.get("supports", []) + if not isinstance(supports, list) or any( + not isinstance(x, str) for x in supports + ): + raise ValueError + passages.append( + Passage(p["id"], p["source"], p["text"], tuple(supports)) + ) + cases.append( + Case(item["id"], item["query"], tuple(required), tuple(passages)) + ) + if not cases or len({case.id for case in cases}) != len(cases): + raise ValueError + return tuple(cases), hashlib.sha256(raw).hexdigest() + except (KeyError, TypeError, ValueError, AttributeError): + raise ValueError("invalid corpus schema or evidence labels") from None + + +async def remote_scores(case: Case, name: str, args) -> dict: + if not args.allow_network: + return {"status": "skipped", "reason": "network_disabled"} + # Dedicated experiment credentials, never app .env files or application globals. + if name == "embeddings": + endpoint = os.getenv("CONTEXT_EMBEDDING_URL") + key = os.getenv("CONTEXT_EMBEDDING_KEY") + model = os.getenv("CONTEXT_EMBEDDING_MODEL") + else: + endpoint = os.getenv("CONTEXT_JEV_URL", "https://api.typesafe.ai/v1/systemone") + key = os.getenv("TYPESAFE_API_KEY") + model = os.getenv("CONTEXT_JEV_MODEL", "jev-latest") + if not endpoint or not key or not model: + return {"status": "skipped", "reason": "missing_provider_configuration"} + import httpx + from .providers import HTTPScorer, ProviderError + + started = time.perf_counter() + scorer = None + try: + async with httpx.AsyncClient() as client: + scorer = HTTPScorer( + client, + endpoint, + model, + key, + concurrency=args.concurrency, + timeout=args.timeout, + ) + result = await getattr(scorer, name)(case.query, case.passages) + return { + "status": "ok", + **asdict(result), + "elapsed_ms": (time.perf_counter() - started) * 1000, + } + except ProviderError as error: + return { + "status": "error", + "reason": error.code, + "requests": scorer.requests if scorer else 0, + "reported_input_tokens": scorer.reported_input_tokens if scorer else 0, + "input_tokens": None, + } + except TimeoutError: + return { + "status": "error", + "reason": "deadline_exceeded", + "requests": scorer.requests if scorer else 0, + "reported_input_tokens": scorer.reported_input_tokens if scorer else 0, + "input_tokens": None, + } + except (ValueError, httpx.HTTPError): + return {"status": "error", "reason": "invalid_provider_configuration"} + + +def summaries(rows: list[dict]) -> list[dict]: + result = [] + for strategy in dict.fromkeys(row["strategy"] for row in rows): + group = [row for row in rows if row["strategy"] == strategy] + ok = [row for row in group if row["status"] == "ok"] + + def average(field): + values = [row[field] for row in ok if row[field] is not None] + return sum(values) / len(values) if values else None + + result.append( + { + "strategy": strategy, + "completed": len(ok), + "precision_cases": sum( + row["evidence_precision"] is not None for row in ok + ), + "recall_cases": sum(row["evidence_recall"] is not None for row in ok), + "skipped": sum(row["status"] == "skipped" for row in group), + "errors": sum(row["status"] == "error" for row in group), + "macro_precision": average("evidence_precision"), + "macro_evidence_recall": average("evidence_recall"), + "mean_selected_chars": average("selected_chars"), + "answerable_cases_missing_evidence": sum( + bool(row["missing_evidence"]) for row in ok + ), + "correct_abstentions": sum( + row["correct_abstention"] is True for row in ok + ), + } + ) + return result + + +async def replay(args) -> dict: + cases, fingerprint = load_cases(args.corpus) + budget = Budget(args.max_results, args.max_chars) + if len(args.scorers) != len(set(args.scorers)): + raise ValueError("scorers must be unique") + if not 0 <= number(args.relative) <= 1 or not 0 <= number(args.jev_min) <= 3: + raise ValueError("invalid selection threshold") + if not -1 <= number(args.cosine_min) <= 1: + raise ValueError("cosine threshold must be in [-1, 1]") + if args.concurrency < 1 or number(args.timeout) <= 0: + raise ValueError("concurrency and timeout must be positive") + if args.max_candidates < 1 or args.max_input_chars < 1: + raise ValueError("input limits must be positive") + if any( + len(case.passages) > args.max_candidates + or len(case.query) + sum(len(p.text) for p in case.passages) + > args.max_input_chars + for case in cases + ): + raise ValueError("corpus exceeds per-case input limits; no provider was called") + cached = {} + if args.reuse_scores: + saved = json.loads(args.reuse_scores.read_text()) + if ( + saved.get("corpus_sha256") != fingerprint + or saved.get("schema_version") != 1 + ): + raise ValueError("saved scores belong to a different corpus/schema") + if not isinstance(saved.get("score_runs"), list): + raise ValueError("invalid saved score runs") + pool_sizes = {case.id: len(case.passages) for case in cases} + for run in saved["score_runs"]: + if run.get("status") not in ("ok", "skipped", "error"): + raise ValueError("invalid saved status") + if run["case"] not in pool_sizes or run["scorer"] not in ( + "bm25", + "embeddings", + "jev", + ): + raise ValueError("invalid saved scoring identity") + if run["status"] == "ok": + values = run.get("values") + if ( + not isinstance(values, list) + or len(values) != pool_sizes[run["case"]] + ): + raise ValueError("invalid saved score count") + values = [number(value) for value in values] + if run["scorer"] == "jev" and any( + not 0 <= value <= 3 for value in values + ): + raise ValueError("invalid saved usefulness") + if run["scorer"] == "embeddings" and any( + not -1 <= value <= 1 for value in values + ): + raise ValueError("invalid saved cosine") + key = (run["case"], run["scorer"]) + if key in cached: + raise ValueError("duplicate saved scoring run") + cached[key] = run + rows, score_runs = [], [] + for case in cases: + for strategy, cap in [ + ("all", len(case.passages)), + ("input-topk", budget.max_results), + ]: + rows.append( + { + "case": case.id, + "strategy": strategy, + "status": "ok", + **measure( + case, select(case.passages, Budget(cap, budget.max_chars)) + ), + } + ) + for name in args.scorers: + if args.reuse_scores and name != "bm25": + run = dict( + cached.get( + (case.id, name), + {"status": "skipped", "reason": "no_saved_scores"}, + ) + ) + run.update({"requests": 0, "input_tokens": 0, "origin": "saved_scores"}) + elif name == "bm25": + started = time.perf_counter() + run = { + "status": "ok", + "values": bm25(case.query, case.passages), + "model": "bm25-k1=1.5-b=0.75-nfkc", + "requests": 0, + "input_tokens": 0, + "elapsed_ms": (time.perf_counter() - started) * 1000, + } + else: + run = await remote_scores(case, name, args) + run.update({"case": case.id, "scorer": name}) + score_runs.append(run) + policies = [(f"{name}-topk", {})] + policies.append( + ( + f"{name}-filter", + ( + {"relative": args.relative} + if name == "bm25" + else { + "minimum": ( + args.jev_min if name == "jev" else args.cosine_min + ) + } + ), + ) + ) + for strategy, threshold in policies: + row = {"case": case.id, "strategy": strategy, "status": run["status"]} + if run["status"] == "ok": + selected = select( + case.passages, + budget, + tuple(run["values"]), + deduplicate=args.deduplicate, + **threshold, + ) + row.update(measure(case, selected)) + else: + row["reason"] = run["reason"] + rows.append(row) + return { + "schema_version": 1, + "implementation_sha256": hashlib.sha256( + b"".join( + Path(__file__).with_name(name).read_bytes() + for name in ("core.py", "providers.py", "replay.py") + ) + ).hexdigest(), + "python_version": platform.python_version(), + "corpus_sha256": fingerprint, + "config": { + "max_results": budget.max_results, + "max_chars": budget.max_chars, + "deduplicate": args.deduplicate, + "relative": args.relative, + "jev_min": args.jev_min, + "cosine_min": args.cosine_min, + "concurrency": args.concurrency, + "timeout": args.timeout, + "max_candidates": args.max_candidates, + "max_input_chars": args.max_input_chars, + }, + "score_runs": score_runs, + "rows": rows, + "summary": summaries(rows), + } + + +def parser() -> argparse.ArgumentParser: + result = argparse.ArgumentParser(description=__doc__) + result.add_argument( + "--corpus", type=Path, default=Path(__file__).with_name("cases.json") + ) + result.add_argument("--output", type=Path) + result.add_argument("--reuse-scores", type=Path) + result.add_argument( + "--scorers", + nargs="+", + choices=["bm25", "embeddings", "jev"], + default=["bm25", "embeddings", "jev"], + ) + result.add_argument("--allow-network", action="store_true") + result.add_argument("--deduplicate", action="store_true") + result.add_argument("--max-results", type=int, default=10) + result.add_argument("--max-chars", type=int, default=12000) + result.add_argument("--relative", type=float, default=0.5) + result.add_argument("--jev-min", type=float, default=1.5) + result.add_argument("--cosine-min", type=float, default=0.35) + result.add_argument("--concurrency", type=int, default=4) + result.add_argument("--timeout", type=float, default=30) + result.add_argument("--max-candidates", type=int, default=256) + result.add_argument("--max-input-chars", type=int, default=256000) + return result + + +def main(): + arguments = parser() + args = arguments.parse_args() + try: + report = asyncio.run(replay(args)) + except (ValueError, KeyError, TypeError, OSError): + arguments.error("invalid corpus, saved scores, configuration or output path") + rendered = json.dumps(report, indent=2, allow_nan=False) + "\n" + if args.output: + args.output.write_text(rendered) + print(json.dumps(report["summary"], indent=2)) + else: + print(rendered, end="") + # A failed requested scorer must not be presented as a successful experiment. + if any(row["status"] == "error" for row in report["rows"]): + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/tests/test_context_selection.py b/tests/test_context_selection.py new file mode 100644 index 00000000..febfecfc --- /dev/null +++ b/tests/test_context_selection.py @@ -0,0 +1,479 @@ +"""Synthetic/HTTP-emulated contract checks, not evidence of model quality.""" + +import asyncio +from dataclasses import replace +import json +from pathlib import Path +from unittest.mock import patch + +import httpx +import pytest + +from evals.context_selection.core import ( + Budget, + Case, + Passage, + bm25, + cosine, + measure, + select, +) +from evals.context_selection.providers import HTTPScorer, ProviderError +from evals.context_selection.replay import load_cases, parser, replay + + +P = ( + Passage("a", "source-a", "first fact", ("fact",)), + Passage("b", "source-b", "second text"), +) + + +def test_empty_selection_is_success_not_fallback(): + case = Case("case", "question", ("fact",), P) + result = measure(case, select(P, Budget(), (0.2, 0.3), minimum=1.5)) + assert result["selected_ids"] == [] + assert result["evidence_recall"] == 0 + assert result["evidence_precision"] is None + assert result["missing_evidence"] == ["fact"] + + +def test_stable_ties_preserve_original_provenance(): + chosen = select(P, Budget(2), (2, 2)) + assert chosen == P + assert chosen[0] is P[0] + assert measure(Case("case", "query", ("fact",), P), chosen)["sources"] == [ + "source-a", + "source-b", + ] + + +def test_character_limit_keeps_whole_passages_not_truncated_evidence(): + chosen = select(P, Budget(10, len(P[1].text)), (1, 2)) + assert chosen == (P[1],) + assert select(P, Budget(0), (1, 2)) == () + assert select(P, Budget(10, 0), (1, 2)) == () + + +def test_exact_duplicate_suppression_is_explicit(): + duplicate = replace(P[0], id="c", source="mirror") + assert len(select((P[0], duplicate), Budget(), (1, 1))) == 2 + assert select((P[0], duplicate), Budget(), (1, 1), deduplicate=True) == (P[0],) + + +@pytest.mark.parametrize( + "scores", [(1,), (True, 1), (float("nan"), 1), (float("inf"), 1)] +) +def test_malformed_scores_rejected(scores): + with pytest.raises(ValueError): + select(P, Budget(), scores) + + +@pytest.mark.parametrize("value", [-1, 2, float("nan")]) +def test_invalid_relative_threshold_rejected(value): + with pytest.raises(ValueError): + select(P, Budget(), (1, 2), relative=value) + + +def test_relative_filter_does_not_pad_zero_match(): + scores = bm25("unmatched-term", P) + assert scores == (0, 0) + assert select(P, Budget(), scores, relative=0.5) == () + assert len(select(P, Budget(), scores)) == 2 + + +def test_bm25_preserves_negations_identifiers_and_unicode(): + passages = ( + Passage("a", "s", "BUG_1842 not deleted 保管"), + Passage("b", "s", "unrelated"), + ) + assert bm25("BUG_1842", passages)[0] > 0 + assert bm25("not", passages)[0] > 0 + assert bm25("保管", passages)[0] > 0 + + +@pytest.mark.parametrize( + "left,right", + [([0, 0], [1, 1]), ([1], [1, 1]), ([float("nan")], [1]), ([True], [1])], +) +def test_invalid_vectors_rejected(left, right): + with pytest.raises(ValueError): + cosine(left, right) + + +def test_cosine_handles_non_unit_vectors(): + assert cosine([3, 4], [6, 8]) == pytest.approx(1) + assert cosine([3, 4], [-6, -8]) == pytest.approx(-1) + + +def test_metric_detects_exception_loss_and_rewritten_citation(): + exception = replace(P[1], supports=("exception",)) + case = Case("case", "query", ("fact", "exception"), (P[0], exception)) + result = measure(case, (P[0],)) + assert result["evidence_precision"] == 1 + assert result["evidence_recall"] == 0.5 + assert result["missing_evidence"] == ["exception"] + with pytest.raises(ValueError): + measure(case, (replace(P[0], text="invented"),)) + + +@pytest.mark.parametrize( + "passages,required", + [((P[0], P[0]), ("fact",)), (P, ("absent",)), (P, ("fact", "fact"))], +) +def test_invalid_evidence_corpus_rejected(passages, required): + with pytest.raises(ValueError): + Case("case", "query", required, passages) + + +def scorer(client, **kwargs): + return HTTPScorer( + client, "https://provider.invalid/v1", "model-test", "secret-test", **kwargs + ) + + +def jev_body(score): + return {"answers": {"usefulness": {"score": score}}, "usage": {"input_tokens": 4}} + + +async def test_jev_uses_text_only_not_labels_and_preserves_request_order(): + payloads = [] + + async def handler(request): + payload = json.loads(request.content) + payloads.append(payload) + if payload["state"] == P[0].text: + await asyncio.sleep(0.01) + return httpx.Response( + 200, json=jev_body(2 if payload["state"] == P[0].text else 0) + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + result = await scorer(client).jev("query", P) + assert result.values == (2, 0) + assert result.input_tokens == 8 + assert result.requests == 2 + assert all("supports" not in p and "source" not in p for p in payloads) + + +async def test_concurrency_limit_shared_across_overlapping_calls(): + active = maximum = 0 + + async def handler(request): + nonlocal active, maximum + active += 1 + maximum = max(maximum, active) + try: + await asyncio.sleep(0.01) + return httpx.Response(200, json=jev_body(2)) + finally: + active -= 1 + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + adapter = scorer(client, concurrency=2) + await asyncio.gather(adapter.jev("q1", P * 3), adapter.jev("q2", P * 3)) + assert maximum == 2 + assert active == 0 + + +@pytest.mark.parametrize("score", [-1, 4, float("inf"), "2", True]) +async def test_jev_rejects_malformed_scores(score): + async with httpx.AsyncClient( + transport=httpx.MockTransport( + lambda req: httpx.Response(200, json=jev_body(score)) + ) + ) as client: + with pytest.raises(ProviderError, match="invalid_usefulness"): + await scorer(client).jev("query", P) + + +async def test_failure_cancels_and_drains_siblings_then_retry_works(): + running = asyncio.Event() + stopped = asyncio.Event() + fail = True + + async def handler(request): + if not fail: + return httpx.Response(200, json=jev_body(2)) + payload = json.loads(request.content) + if payload["state"] == P[1].text: + running.set() + try: + await asyncio.Event().wait() + finally: + stopped.set() + await running.wait() + return httpx.Response(401, text="secret-test and private passage") + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + adapter = scorer(client) + with pytest.raises(ProviderError) as caught: + await adapter.jev("query", P) + assert str(caught.value) == "http_401" + assert stopped.is_set() + fail = False + assert (await adapter.jev("query", P)).values == (2, 2) + + +@pytest.mark.parametrize("cancel", [False, True]) +async def test_deadline_and_cancellation_release_active_calls(cancel): + active = 0 + started = asyncio.Event() + + async def handler(request): + nonlocal active + active += 1 + started.set() + try: + await asyncio.Event().wait() + finally: + active -= 1 + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + task = asyncio.create_task(scorer(client, timeout=0.03).jev("query", P)) + await started.wait() + if cancel: + task.cancel() + with pytest.raises(asyncio.CancelledError if cancel else TimeoutError): + await task + assert active == 0 + + +async def test_embedding_response_indices_not_response_order(): + body = { + "data": [ + {"index": 2, "embedding": [0, 1]}, + {"index": 0, "embedding": [2, 0]}, + {"index": 1, "embedding": [3, 0]}, + ], + "usage": {"prompt_tokens": 9}, + } + async with httpx.AsyncClient( + transport=httpx.MockTransport(lambda req: httpx.Response(200, json=body)) + ) as client: + result = await scorer(client).embeddings("query", P) + assert result.values == (1, 0) + assert result.input_tokens == 9 + + +@pytest.mark.parametrize( + "rows", + [ + [], + [{"index": 0, "embedding": [1, 0]}] * 3, + [{"index": i, "embedding": [0, 0]} for i in range(3)], + ], +) +async def test_bad_embedding_batches_rejected(rows): + async with httpx.AsyncClient( + transport=httpx.MockTransport( + lambda req: httpx.Response(200, json={"data": rows}) + ) + ) as client: + with pytest.raises(ProviderError, match="invalid_embeddings"): + await scorer(client).embeddings("query", P) + + +async def test_default_replay_never_constructs_network_client_even_with_credentials(): + args = parser().parse_args([]) + with patch.dict( + "os.environ", + {"TYPESAFE_API_KEY": "secret-test", "CONTEXT_EMBEDDING_KEY": "secret-test"}, + ), patch("httpx.AsyncClient", side_effect=AssertionError("network forbidden")): + report = await replay(args) + assert all( + run["status"] == "skipped" + for run in report["score_runs"] + if run["scorer"] != "bm25" + ) + assert all( + row["status"] == "ok" + for row in report["rows"] + if row["strategy"].startswith("bm25") + ) + assert "secret-test" not in json.dumps(report) + + +async def test_replay_reuses_saved_scores_for_ablation_without_network(tmp_path): + args = parser().parse_args([]) + report = await replay(args) + cases, _ = load_cases(args.corpus) + for case in cases: + report["score_runs"].append( + { + "case": case.id, + "scorer": "jev", + "status": "ok", + "values": [2 if p.supports else 0 for p in case.passages], + "model": "FAKE-NOT-JEV", + "requests": 99, + } + ) + report["score_runs"] = [ + r + for r in report["score_runs"] + if not (r["scorer"] == "jev" and r["status"] == "skipped") + ] + saved = tmp_path / "scores.json" + saved.write_text(json.dumps(report)) + args.reuse_scores = saved + with patch("httpx.AsyncClient", side_effect=AssertionError("network forbidden")): + reused = await replay(args) + filtered = [r for r in reused["rows"] if r["strategy"] == "jev-filter"] + assert all(row["complete_evidence"] for row in filtered) + assert all( + run["requests"] == 0 for run in reused["score_runs"] if run["scorer"] == "jev" + ) + report["corpus_sha256"] = "different" + saved.write_text(json.dumps(report)) + with pytest.raises(ValueError, match="different corpus"): + await replay(args) + + +async def test_input_limits_reject_before_provider_use(): + args = parser().parse_args(["--max-candidates", "1", "--allow-network"]) + with patch( + "httpx.AsyncClient", side_effect=AssertionError("network forbidden") + ), pytest.raises(ValueError, match="input limits"): + await replay(args) + + +async def test_empty_provider_pools_make_no_http_requests(): + async with httpx.AsyncClient( + transport=httpx.MockTransport(lambda req: pytest.fail("unexpected inference")) + ) as client: + adapter = scorer(client) + assert (await adapter.embeddings("query", ())).requests == 0 + assert (await adapter.jev("query", ())).requests == 0 + assert adapter.requests == 0 + + +async def test_jev_records_resolved_model_and_unknown_usage(): + body = {"model": "resolved-model", "answers": {"usefulness": {"score": 2}}} + async with httpx.AsyncClient( + transport=httpx.MockTransport(lambda req: httpx.Response(200, json=body)) + ) as client: + result = await scorer(client).jev("query", P) + assert result.model == "resolved-model" + assert result.input_tokens is None + + +async def test_jev_rejects_mixed_resolved_models(): + async def handler(request): + body = jev_body(2) + body["model"] = json.loads(request.content)["state"] + return httpx.Response(200, json=body) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + with pytest.raises(ProviderError, match="inconsistent_model"): + await scorer(client).jev("query", P) + + +@pytest.mark.parametrize("status", [302, 429, 529]) +async def test_provider_errors_are_not_retried_or_redirected(status): + calls = 0 + + def handler(request): + nonlocal calls + calls += 1 + return httpx.Response( + status, headers={"location": "https://other.invalid"}, text="secret-test" + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + adapter = scorer(client, concurrency=1) + with pytest.raises(ProviderError, match=f"http_{status}"): + await adapter.jev("query", P) + assert adapter.requests == calls == 1 + + +async def test_missing_live_credentials_are_skipped_before_client_creation(): + args = parser().parse_args(["--allow-network"]) + with patch.dict("os.environ", {}, clear=True), patch( + "httpx.AsyncClient", side_effect=AssertionError("network forbidden") + ): + report = await replay(args) + assert all( + run["reason"] == "missing_provider_configuration" + for run in report["score_runs"] + if run["scorer"] != "bm25" + ) + + +@pytest.mark.parametrize("change", ["status", "count", "range", "duplicate"]) +async def test_saved_score_validation_rejects_corruption(tmp_path, change): + args = parser().parse_args([]) + report = await replay(args) + cases, _ = load_cases(args.corpus) + run = { + "case": cases[0].id, + "scorer": "jev", + "status": "ok", + "values": [2] * len(cases[0].passages), + } + report["score_runs"] = [run] + if change == "status": + run["status"] = "unknown" + elif change == "count": + run["values"] = [] + elif change == "range": + run["values"][0] = 4 + else: + report["score_runs"].append(run) + saved = tmp_path / "corrupt.json" + saved.write_text(json.dumps(report)) + args.reuse_scores = saved + with pytest.raises(ValueError): + await replay(args) + + +def test_cli_reports_failed_provider_without_claiming_success( + monkeypatch, tmp_path, capsys +): + from evals.context_selection import replay as module + + async def failed(*args): + return { + "status": "error", + "reason": "http_429", + "requests": 1, + "input_tokens": None, + } + + monkeypatch.setattr(module, "remote_scores", failed) + output = tmp_path / "report.json" + monkeypatch.setattr( + "sys.argv", ["replay", "--scorers", "jev", "--output", str(output)] + ) + with pytest.raises(SystemExit) as caught: + module.main() + assert caught.value.code == 1 + report = json.loads(output.read_text()) + assert all( + row["status"] == "error" + for row in report["rows"] + if row["strategy"].startswith("jev") + ) + assert "secret-test" not in capsys.readouterr().out + + +async def test_failed_batch_exposes_attempt_counts_not_provider_body(monkeypatch): + from evals.context_selection import replay as module + + class Client(httpx.AsyncClient): + def __init__(self): + super().__init__( + transport=httpx.MockTransport( + lambda req: httpx.Response(429, text="secret-test") + ) + ) + + args = parser().parse_args(["--allow-network", "--concurrency", "1"]) + monkeypatch.setenv("TYPESAFE_API_KEY", "secret-test") + monkeypatch.setattr("httpx.AsyncClient", Client) + result = await module.remote_scores( + Case("case", "query", ("fact",), P), "jev", args + ) + assert result["status"] == "error" + assert result["requests"] == 1 + assert result["input_tokens"] is None + assert "secret-test" not in json.dumps(result)