retrieval-eval-gate 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- rag_eval_gate/__init__.py +39 -0
- rag_eval_gate/audit.py +257 -0
- rag_eval_gate/backends.py +253 -0
- rag_eval_gate/cli.py +403 -0
- rag_eval_gate/conformal.py +444 -0
- rag_eval_gate/goldset.py +136 -0
- rag_eval_gate/harness.py +281 -0
- rag_eval_gate/metrics.py +106 -0
- rag_eval_gate/power.py +154 -0
- rag_eval_gate/ppi.py +224 -0
- rag_eval_gate/report.py +160 -0
- rag_eval_gate/sequential.py +312 -0
- rag_eval_gate/stats.py +415 -0
- rag_eval_gate/version.py +3 -0
- retrieval_eval_gate-0.1.0.dist-info/METADATA +376 -0
- retrieval_eval_gate-0.1.0.dist-info/RECORD +20 -0
- retrieval_eval_gate-0.1.0.dist-info/WHEEL +5 -0
- retrieval_eval_gate-0.1.0.dist-info/entry_points.txt +3 -0
- retrieval_eval_gate-0.1.0.dist-info/licenses/LICENSE +21 -0
- retrieval_eval_gate-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
"""Retrieval evaluation that reports what it can and cannot resolve.
|
|
2
|
+
|
|
3
|
+
A score is an estimate from a finite query sample, so nothing here reports a bare number.
|
|
4
|
+
Metrics arrive with intervals, comparisons with p-values corrected for having asked several
|
|
5
|
+
questions at once, and every gold set with a statement of the differences it is too small
|
|
6
|
+
to detect.
|
|
7
|
+
|
|
8
|
+
from rag_eval_gate import HttpBackend, load_gold, evaluate, add_intervals
|
|
9
|
+
|
|
10
|
+
queries = load_gold("gold.jsonl")
|
|
11
|
+
backend = HttpBackend("http://localhost:8080/api/search?q={query}&mode={mode}")
|
|
12
|
+
result = evaluate(backend, queries, "hybrid")
|
|
13
|
+
add_intervals({"hybrid": result}, iters=10_000)
|
|
14
|
+
print(result.means["mrr@10"], result.intervals["mrr@10"])
|
|
15
|
+
|
|
16
|
+
The statistical layer is usable on its own — `stats` for intervals and paired tests,
|
|
17
|
+
`sequential` for anytime-valid gates, `conformal` for distribution-free guarantees, and
|
|
18
|
+
`ppi` for a valid interval on a quantity only a biased LLM judge has scored at scale.
|
|
19
|
+
"""
|
|
20
|
+
from .audit import (ComparisonAudit, MetricAudit, audit_comparison, audit_metric,
|
|
21
|
+
differing_query_bounds)
|
|
22
|
+
from .backends import (Backend, BackendError, CommandBackend, Hit, HttpBackend,
|
|
23
|
+
RunFileBackend, build_backend)
|
|
24
|
+
from .goldset import GoldQuery, GoldSetError, load_gold, shape
|
|
25
|
+
from .harness import (EvalRun, ModeResult, add_intervals, compare_modes, design_of,
|
|
26
|
+
evaluate, evaluate_sequential)
|
|
27
|
+
from .metrics import METRICS, QueryScore, aggregate, score_query, values_of
|
|
28
|
+
from .ppi import PPIEstimate, classical_interval, ppi_mean
|
|
29
|
+
from .version import __version__
|
|
30
|
+
|
|
31
|
+
__all__ = [
|
|
32
|
+
"MetricAudit", "ComparisonAudit", "audit_metric", "audit_comparison",
|
|
33
|
+
"differing_query_bounds",
|
|
34
|
+
"Backend", "BackendError", "CommandBackend", "Hit", "HttpBackend", "RunFileBackend",
|
|
35
|
+
"build_backend", "GoldQuery", "GoldSetError", "load_gold", "shape",
|
|
36
|
+
"EvalRun", "ModeResult", "add_intervals", "compare_modes", "design_of", "evaluate",
|
|
37
|
+
"evaluate_sequential", "METRICS", "QueryScore", "aggregate", "score_query",
|
|
38
|
+
"values_of", "PPIEstimate", "ppi_mean", "classical_interval", "__version__",
|
|
39
|
+
]
|
rag_eval_gate/audit.py
ADDED
|
@@ -0,0 +1,257 @@
|
|
|
1
|
+
"""Audit a published number without the data behind it.
|
|
2
|
+
|
|
3
|
+
Every retrieval paper, blog post and README has a table like
|
|
4
|
+
|
|
5
|
+
| mode | Recall@5 | MRR@10 |
|
|
6
|
+
| bm25 | 0.70 | 0.67 |
|
|
7
|
+
| hybrid | 1.00 | 0.95 |
|
|
8
|
+
|
|
9
|
+
and almost none of them say how many queries it took or what the numbers can support. That
|
|
10
|
+
is usually answerable anyway: a mean over `n` items, on a metric bounded in [0, 1], already
|
|
11
|
+
pins down more than most authors realise.
|
|
12
|
+
|
|
13
|
+
Three questions get answered from the table alone:
|
|
14
|
+
|
|
15
|
+
1. **How wide is that number really?** With one relevant document per query, Recall@k is a
|
|
16
|
+
count of successes and the exact binomial interval applies — no other information needed.
|
|
17
|
+
For a graded metric the variance is unknown but bounded: a [0, 1] variable with mean `m`
|
|
18
|
+
has variance at most `m(1-m)`, so there is a widest-possible interval and it is worth
|
|
19
|
+
knowing before quoting three decimal places.
|
|
20
|
+
|
|
21
|
+
2. **Could the reported lift have been significant?** A paired randomization test on `k`
|
|
22
|
+
differing queries cannot return a two-sided p below `2/2^k`, and the two reported means
|
|
23
|
+
bound `k` without needing the per-query data. When even the most favourable overlap
|
|
24
|
+
leaves that floor above 0.05, the comparison could not have reached significance — which
|
|
25
|
+
is a fact about the table, not a criticism of the system.
|
|
26
|
+
|
|
27
|
+
3. **What would it take?** The queries needed to resolve an improvement of a given size.
|
|
28
|
+
|
|
29
|
+
Nothing here is a substitute for the per-query data; it is what can be said when the
|
|
30
|
+
per-query data is not on offer, which is the usual case when reading someone else's results.
|
|
31
|
+
"""
|
|
32
|
+
from __future__ import annotations
|
|
33
|
+
|
|
34
|
+
import math
|
|
35
|
+
from dataclasses import dataclass, field
|
|
36
|
+
|
|
37
|
+
from . import stats
|
|
38
|
+
|
|
39
|
+
__all__ = ["MetricAudit", "ComparisonAudit", "audit_metric", "audit_comparison",
|
|
40
|
+
"differing_query_bounds", "render"]
|
|
41
|
+
|
|
42
|
+
ALPHA = 0.05
|
|
43
|
+
POWER = 0.80
|
|
44
|
+
# A reported mean that is not a clean multiple of 1/n cannot be a count of successes.
|
|
45
|
+
BINARY_TOLERANCE = 1e-6
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
@dataclass
|
|
49
|
+
class MetricAudit:
|
|
50
|
+
metric: str
|
|
51
|
+
value: float
|
|
52
|
+
n: int
|
|
53
|
+
binary: bool
|
|
54
|
+
lo: float
|
|
55
|
+
hi: float
|
|
56
|
+
method: str
|
|
57
|
+
successes: int | None = None
|
|
58
|
+
notes: list[str] = field(default_factory=list)
|
|
59
|
+
|
|
60
|
+
@property
|
|
61
|
+
def width(self) -> float:
|
|
62
|
+
return self.hi - self.lo
|
|
63
|
+
|
|
64
|
+
def as_dict(self) -> dict:
|
|
65
|
+
return {"metric": self.metric, "value": self.value, "n": self.n,
|
|
66
|
+
"binary": self.binary, "lo": self.lo, "hi": self.hi,
|
|
67
|
+
"width": self.width, "method": self.method,
|
|
68
|
+
"successes": self.successes, "notes": self.notes}
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
@dataclass
|
|
72
|
+
class ComparisonAudit:
|
|
73
|
+
metric: str
|
|
74
|
+
value: float
|
|
75
|
+
baseline: float
|
|
76
|
+
n: int
|
|
77
|
+
binary: bool
|
|
78
|
+
delta: float
|
|
79
|
+
k_min: int
|
|
80
|
+
k_max: int
|
|
81
|
+
floor_p_best: float # the smallest p attainable, under the kindest overlap
|
|
82
|
+
floor_p_worst: float # ... and under the least kind
|
|
83
|
+
resolvable: bool # can *any* overlap reach p < 0.05?
|
|
84
|
+
determined: bool # do the two means pin k down exactly?
|
|
85
|
+
queries_for_delta: int | None
|
|
86
|
+
notes: list[str] = field(default_factory=list)
|
|
87
|
+
|
|
88
|
+
def as_dict(self) -> dict:
|
|
89
|
+
return {"metric": self.metric, "value": self.value, "baseline": self.baseline,
|
|
90
|
+
"n": self.n, "binary": self.binary, "delta": self.delta,
|
|
91
|
+
"k_min": self.k_min, "k_max": self.k_max,
|
|
92
|
+
"floor_p_best": self.floor_p_best, "floor_p_worst": self.floor_p_worst,
|
|
93
|
+
"resolvable": self.resolvable, "determined": self.determined,
|
|
94
|
+
"queries_for_delta": self.queries_for_delta, "notes": self.notes}
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def _looks_binary(value: float, n: int) -> bool:
|
|
98
|
+
"""True when the mean is a clean count of successes over n."""
|
|
99
|
+
return abs(value * n - round(value * n)) < BINARY_TOLERANCE * max(1, n)
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def audit_metric(metric: str, value: float, n: int, binary: bool | None = None) -> MetricAudit:
|
|
103
|
+
"""The interval a single reported mean actually supports.
|
|
104
|
+
|
|
105
|
+
`binary=None` infers it: Recall@k with one relevant document per query is a per-query
|
|
106
|
+
0/1 outcome, and its mean is then a multiple of 1/n. The inference is reported rather
|
|
107
|
+
than assumed silently, because it changes which interval applies.
|
|
108
|
+
"""
|
|
109
|
+
if not 0.0 <= value <= 1.0:
|
|
110
|
+
raise ValueError(f"{metric} must lie in [0, 1], got {value}")
|
|
111
|
+
if n < 2:
|
|
112
|
+
raise ValueError(f"need at least 2 queries, got {n}")
|
|
113
|
+
|
|
114
|
+
notes: list[str] = []
|
|
115
|
+
if binary is None:
|
|
116
|
+
binary = metric.lower().startswith("recall") and _looks_binary(value, n)
|
|
117
|
+
if binary:
|
|
118
|
+
notes.append("treated as a per-query 0/1 outcome (Recall@k with one relevant "
|
|
119
|
+
"document per query); pass --graded if that is wrong")
|
|
120
|
+
|
|
121
|
+
if binary:
|
|
122
|
+
if not _looks_binary(value, n):
|
|
123
|
+
raise ValueError(f"{value} is not a multiple of 1/{n}, so it cannot be a "
|
|
124
|
+
f"success count over {n} queries")
|
|
125
|
+
successes = round(value * n)
|
|
126
|
+
lo, hi = stats.clopper_pearson(successes, n, ALPHA)
|
|
127
|
+
return MetricAudit(metric, value, n, True, lo, hi, "exact binomial",
|
|
128
|
+
successes, notes)
|
|
129
|
+
|
|
130
|
+
# Graded metric: the per-query spread is unknown, but a [0, 1] variable with mean m has
|
|
131
|
+
# variance at most m(1-m) (attained by a Bernoulli). That is the widest the interval can
|
|
132
|
+
# be — the conservative thing to quote when the per-query scores are not on offer.
|
|
133
|
+
half = stats.NormalDist().inv_cdf(1 - ALPHA / 2) * math.sqrt(value * (1 - value) / n)
|
|
134
|
+
notes.append("per-query scores unknown; this is the widest the interval can be, "
|
|
135
|
+
"attained when every query scores 0 or 1")
|
|
136
|
+
return MetricAudit(metric, value, n, False, max(0.0, value - half),
|
|
137
|
+
min(1.0, value + half), "maximum-variance bound", None, notes)
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def differing_query_bounds(value: float, baseline: float, n: int,
|
|
141
|
+
binary: bool) -> tuple[int, int, bool]:
|
|
142
|
+
"""How many queries must differ between two systems, given only their means.
|
|
143
|
+
|
|
144
|
+
Returns (minimum, maximum, exactly determined).
|
|
145
|
+
|
|
146
|
+
For a 0/1 metric with `a` and `b` successes out of `n`, write `d1` for the queries the
|
|
147
|
+
first system wins and `d2` for those it loses. Then `d1 - d2 = a - b` and the count that
|
|
148
|
+
differs is `d1 + d2`, with `d1` ranging over `[max(0, a-b), min(a, n-b)]`. When that
|
|
149
|
+
range collapses to a point — which is what happens against a perfect score — the two
|
|
150
|
+
published means determine the number exactly.
|
|
151
|
+
|
|
152
|
+
For a graded metric each differing query can move the mean by at most `1/n`, so at least
|
|
153
|
+
`ceil(|delta| * n)` of them differ; the upper bound is `n`.
|
|
154
|
+
"""
|
|
155
|
+
delta = value - baseline
|
|
156
|
+
if binary:
|
|
157
|
+
a, b = round(value * n), round(baseline * n)
|
|
158
|
+
low_d1, high_d1 = max(0, a - b), min(a, n - b)
|
|
159
|
+
k_min = 2 * low_d1 - (a - b)
|
|
160
|
+
k_max = 2 * high_d1 - (a - b)
|
|
161
|
+
return k_min, k_max, k_min == k_max
|
|
162
|
+
return math.ceil(abs(delta) * n), n, False
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def audit_comparison(metric: str, value: float, baseline: float, n: int,
|
|
166
|
+
binary: bool | None = None) -> ComparisonAudit:
|
|
167
|
+
"""What a reported lift between two means can and cannot establish."""
|
|
168
|
+
if binary is None:
|
|
169
|
+
binary = (metric.lower().startswith("recall")
|
|
170
|
+
and _looks_binary(value, n) and _looks_binary(baseline, n))
|
|
171
|
+
for name, v in (("value", value), ("baseline", baseline)):
|
|
172
|
+
if not 0.0 <= v <= 1.0:
|
|
173
|
+
raise ValueError(f"{name} must lie in [0, 1], got {v}")
|
|
174
|
+
if n < 2:
|
|
175
|
+
raise ValueError(f"need at least 2 queries, got {n}")
|
|
176
|
+
|
|
177
|
+
k_min, k_max, determined = differing_query_bounds(value, baseline, n, binary)
|
|
178
|
+
# More differing queries means a smaller attainable p, so k_max gives the best case.
|
|
179
|
+
floor_best = stats.min_attainable_p(k_max)
|
|
180
|
+
floor_worst = stats.min_attainable_p(k_min)
|
|
181
|
+
|
|
182
|
+
notes: list[str] = []
|
|
183
|
+
if determined:
|
|
184
|
+
notes.append(f"the two means determine it exactly: {k_min} of {n} queries differ")
|
|
185
|
+
elif binary:
|
|
186
|
+
notes.append(f"between {k_min} and {k_max} of {n} queries differ, depending on "
|
|
187
|
+
f"which ones each system got right")
|
|
188
|
+
else:
|
|
189
|
+
notes.append(f"at least {k_min} of {n} queries differ (each one can move the mean "
|
|
190
|
+
f"by at most 1/{n}); the per-query data would pin it down")
|
|
191
|
+
|
|
192
|
+
# Sample size to resolve this difference again. A paired test cares about the spread
|
|
193
|
+
# of the per-query *differences*, not of the raw scores. For a 0/1 metric that spread is
|
|
194
|
+
# exact given k: each difference is -1, 0 or +1, so Var(d) = k/n - delta^2. Taking the
|
|
195
|
+
# largest admissible k gives the conservative answer. For a graded metric the
|
|
196
|
+
# differences are unknown and the [0, 1] variance bound is the best available.
|
|
197
|
+
delta = value - baseline
|
|
198
|
+
if binary:
|
|
199
|
+
spread = math.sqrt(max(0.0, k_max / n - delta ** 2))
|
|
200
|
+
else:
|
|
201
|
+
spread = math.sqrt(max(value * (1 - value), baseline * (1 - baseline)))
|
|
202
|
+
needed = stats.required_queries(abs(delta), spread, ALPHA, POWER)
|
|
203
|
+
|
|
204
|
+
return ComparisonAudit(metric, value, baseline, n, binary, value - baseline,
|
|
205
|
+
k_min, k_max, floor_best, floor_worst,
|
|
206
|
+
resolvable=floor_best < ALPHA, determined=determined,
|
|
207
|
+
queries_for_delta=needed, notes=notes)
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
def render(metric_audit: MetricAudit, comparison: ComparisonAudit | None) -> str:
|
|
211
|
+
out: list[str] = []
|
|
212
|
+
add = out.append
|
|
213
|
+
m = metric_audit
|
|
214
|
+
|
|
215
|
+
add(f"{m.metric} = {m.value:.3f} on {m.n} queries")
|
|
216
|
+
if m.successes is not None:
|
|
217
|
+
add(f" that is {m.successes}/{m.n} successes")
|
|
218
|
+
for note in m.notes:
|
|
219
|
+
add(f" note: {note}")
|
|
220
|
+
add("")
|
|
221
|
+
add(f" 95% interval [{m.lo:.3f}, {m.hi:.3f}] width {m.width:.3f} ({m.method})")
|
|
222
|
+
if m.width > 0.15:
|
|
223
|
+
add(f" The reported figure is {m.width * 100:.0f} points wide. Quoting it to two "
|
|
224
|
+
f"decimal places implies a precision {m.n} queries cannot deliver.")
|
|
225
|
+
add("")
|
|
226
|
+
|
|
227
|
+
if comparison is None:
|
|
228
|
+
return "\n".join(out) + "\n"
|
|
229
|
+
|
|
230
|
+
c = comparison
|
|
231
|
+
add(f"Claimed lift {c.delta:+.3f} over {c.baseline:.3f}")
|
|
232
|
+
for note in c.notes:
|
|
233
|
+
add(f" note: {note}")
|
|
234
|
+
add("")
|
|
235
|
+
if c.determined:
|
|
236
|
+
add(f" smallest attainable p {c.floor_p_worst:.4f}")
|
|
237
|
+
else:
|
|
238
|
+
add(f" smallest attainable p {c.floor_p_best:.4f} (best case) to "
|
|
239
|
+
f"{c.floor_p_worst:.4f} (worst)")
|
|
240
|
+
add("")
|
|
241
|
+
if not c.resolvable:
|
|
242
|
+
add(" UNRESOLVABLE. Under every possible overlap of the two systems' successes,")
|
|
243
|
+
add(" a paired test on this many differing queries has a floor above 0.05, so no")
|
|
244
|
+
add(" outcome of that comparison could have reached significance. This is a fact")
|
|
245
|
+
add(" about the sample size, not about the systems.")
|
|
246
|
+
elif c.floor_p_worst >= ALPHA:
|
|
247
|
+
add(" POSSIBLY UNRESOLVABLE. It depends on which queries each system got right —")
|
|
248
|
+
add(" information the published table does not contain. Ask for the per-query")
|
|
249
|
+
add(" scores, or for the paired test.")
|
|
250
|
+
else:
|
|
251
|
+
add(" Resolvable in principle. Whether it *was* resolved needs the paired test;")
|
|
252
|
+
add(" a difference in means is not evidence on its own.")
|
|
253
|
+
add("")
|
|
254
|
+
if c.queries_for_delta:
|
|
255
|
+
add(f" To resolve {abs(c.delta):+.3f} again at 80% power, assuming the widest")
|
|
256
|
+
add(f" per-query spread consistent with these means: {c.queries_for_delta:,} queries.")
|
|
257
|
+
return "\n".join(out) + "\n"
|
|
@@ -0,0 +1,253 @@
|
|
|
1
|
+
"""Where rankings come from — HTTP, a subprocess, or a file already on disk.
|
|
2
|
+
|
|
3
|
+
The evaluation, the statistics and the gate do not care how a ranking was produced, and
|
|
4
|
+
tying them to one API shape is what keeps most eval harnesses locked to the repository
|
|
5
|
+
they were born in. Three adapters cover essentially everything:
|
|
6
|
+
|
|
7
|
+
* `HttpBackend` — a URL template plus a small field map, for a service that answers JSON.
|
|
8
|
+
No assumption about the route, the parameter names, or where the hits sit in the body.
|
|
9
|
+
* `CommandBackend` — runs any command and reads JSON from its stdout. A retrieval system
|
|
10
|
+
written in Go, Rust, TypeScript or a notebook is one shell line away from being gated.
|
|
11
|
+
* `RunFileBackend` — a TREC run file, the format the IR community already exchanges. Rank
|
|
12
|
+
a corpus once, evaluate it many times, with no service running at all.
|
|
13
|
+
|
|
14
|
+
All three return the same thing: `Hit` objects in rank order.
|
|
15
|
+
"""
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import json
|
|
19
|
+
import shlex
|
|
20
|
+
import subprocess
|
|
21
|
+
import time
|
|
22
|
+
import urllib.error
|
|
23
|
+
import urllib.parse
|
|
24
|
+
import urllib.request
|
|
25
|
+
from dataclasses import dataclass, field
|
|
26
|
+
from pathlib import Path
|
|
27
|
+
|
|
28
|
+
__all__ = ["Hit", "Backend", "HttpBackend", "CommandBackend", "RunFileBackend",
|
|
29
|
+
"BackendError", "build_backend"]
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class BackendError(RuntimeError):
|
|
33
|
+
"""A ranking could not be obtained — distinct from a query legitimately scoring zero."""
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
@dataclass(frozen=True)
|
|
37
|
+
class Hit:
|
|
38
|
+
"""One retrieved item. `score` is optional; only calibration needs it."""
|
|
39
|
+
|
|
40
|
+
doc_id: str
|
|
41
|
+
score: float | None = None
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class Backend:
|
|
45
|
+
"""Anything that can rank documents for a query."""
|
|
46
|
+
|
|
47
|
+
def search(self, query: str, mode: str) -> list[Hit]: # pragma: no cover - interface
|
|
48
|
+
raise NotImplementedError
|
|
49
|
+
|
|
50
|
+
def close(self) -> None:
|
|
51
|
+
pass
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _dig(payload, path: str):
|
|
55
|
+
"""Follow a dotted path into nested dicts/lists — `results`, `data.hits`, `hits.hits`."""
|
|
56
|
+
current = payload
|
|
57
|
+
for part in path.split("."):
|
|
58
|
+
if part == "":
|
|
59
|
+
continue
|
|
60
|
+
if isinstance(current, list):
|
|
61
|
+
try:
|
|
62
|
+
current = current[int(part)]
|
|
63
|
+
continue
|
|
64
|
+
except (ValueError, IndexError) as e:
|
|
65
|
+
raise BackendError(f"cannot index list with '{part}' in path '{path}'") from e
|
|
66
|
+
if not isinstance(current, dict) or part not in current:
|
|
67
|
+
available = list(current)[:8] if isinstance(current, dict) else type(current).__name__
|
|
68
|
+
raise BackendError(f"path '{path}': no '{part}' — found {available}")
|
|
69
|
+
current = current[part]
|
|
70
|
+
return current
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
@dataclass
|
|
74
|
+
class HttpBackend(Backend):
|
|
75
|
+
"""A search service that answers JSON.
|
|
76
|
+
|
|
77
|
+
`url_template` is formatted with `query` (already percent-encoded) and `mode`, so any
|
|
78
|
+
route shape works:
|
|
79
|
+
|
|
80
|
+
http://localhost:8080/api/search?q={query}&mode={mode}
|
|
81
|
+
https://api.example.com/v2/retrieve?text={query}&strategy={mode}&k=10
|
|
82
|
+
|
|
83
|
+
`results_path` locates the hit list in the response body and `id_field` / `score_field`
|
|
84
|
+
name the keys inside a hit, so no particular schema is assumed. A hit may also be a
|
|
85
|
+
bare string, in which case it is the document id.
|
|
86
|
+
"""
|
|
87
|
+
|
|
88
|
+
url_template: str
|
|
89
|
+
results_path: str = "results"
|
|
90
|
+
id_field: str = "docId"
|
|
91
|
+
score_field: str | None = "score"
|
|
92
|
+
timeout: float = 120.0
|
|
93
|
+
attempts: int = 3
|
|
94
|
+
headers: dict[str, str] = field(default_factory=dict)
|
|
95
|
+
|
|
96
|
+
def search(self, query: str, mode: str) -> list[Hit]:
|
|
97
|
+
url = self.url_template.format(query=urllib.parse.quote(query), mode=mode)
|
|
98
|
+
payload = self._get(url)
|
|
99
|
+
rows = _dig(payload, self.results_path)
|
|
100
|
+
if not isinstance(rows, list):
|
|
101
|
+
raise BackendError(f"'{self.results_path}' is {type(rows).__name__}, not a list")
|
|
102
|
+
return [self._hit(row, url) for row in rows]
|
|
103
|
+
|
|
104
|
+
def _hit(self, row, url: str) -> Hit:
|
|
105
|
+
if isinstance(row, str):
|
|
106
|
+
return Hit(row)
|
|
107
|
+
if not isinstance(row, dict):
|
|
108
|
+
raise BackendError(f"hit is {type(row).__name__}, expected object or string")
|
|
109
|
+
if self.id_field not in row:
|
|
110
|
+
raise BackendError(f"hit has no '{self.id_field}' — keys: {list(row)[:8]} ({url})")
|
|
111
|
+
score = row.get(self.score_field) if self.score_field else None
|
|
112
|
+
return Hit(str(row[self.id_field]), float(score) if score is not None else None)
|
|
113
|
+
|
|
114
|
+
def _get(self, url: str):
|
|
115
|
+
last: Exception | None = None
|
|
116
|
+
for attempt in range(self.attempts):
|
|
117
|
+
try:
|
|
118
|
+
request = urllib.request.Request(
|
|
119
|
+
url, headers={"User-Agent": "rag-eval-gate", **self.headers})
|
|
120
|
+
with urllib.request.urlopen(request, timeout=self.timeout) as response:
|
|
121
|
+
return json.load(response)
|
|
122
|
+
except (urllib.error.URLError, TimeoutError, json.JSONDecodeError) as e:
|
|
123
|
+
last = e
|
|
124
|
+
if attempt + 1 < self.attempts:
|
|
125
|
+
time.sleep(2 ** attempt)
|
|
126
|
+
raise BackendError(f"{url}: failed after {self.attempts} attempts ({last})") from last
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
@dataclass
|
|
130
|
+
class CommandBackend(Backend):
|
|
131
|
+
"""Runs a command and reads JSON from its stdout.
|
|
132
|
+
|
|
133
|
+
The escape hatch that makes the language of your retrieval system irrelevant:
|
|
134
|
+
|
|
135
|
+
rag-eval-gate eval gold.jsonl \\
|
|
136
|
+
--backend command \\
|
|
137
|
+
--command './my-retriever --query {query} --mode {mode} --json'
|
|
138
|
+
|
|
139
|
+
`{query}` and `{mode}` are substituted as single shell-quoted arguments, so a query
|
|
140
|
+
containing spaces, quotes or a semicolon is passed through as one argument rather than
|
|
141
|
+
reinterpreted by the shell.
|
|
142
|
+
"""
|
|
143
|
+
|
|
144
|
+
command_template: str
|
|
145
|
+
results_path: str = "results"
|
|
146
|
+
id_field: str = "docId"
|
|
147
|
+
score_field: str | None = "score"
|
|
148
|
+
timeout: float = 120.0
|
|
149
|
+
|
|
150
|
+
def search(self, query: str, mode: str) -> list[Hit]:
|
|
151
|
+
command = self.command_template.format(query=shlex.quote(query),
|
|
152
|
+
mode=shlex.quote(mode))
|
|
153
|
+
try:
|
|
154
|
+
finished = subprocess.run(shlex.split(command), capture_output=True, text=True,
|
|
155
|
+
timeout=self.timeout, check=False)
|
|
156
|
+
except (OSError, subprocess.TimeoutExpired) as e:
|
|
157
|
+
raise BackendError(f"{command}: {e}") from e
|
|
158
|
+
if finished.returncode != 0:
|
|
159
|
+
tail = (finished.stderr or "").strip()[-400:]
|
|
160
|
+
raise BackendError(f"exit {finished.returncode}: {command}\n{tail}")
|
|
161
|
+
try:
|
|
162
|
+
payload = json.loads(finished.stdout)
|
|
163
|
+
except json.JSONDecodeError as e:
|
|
164
|
+
raise BackendError(f"{command}: stdout is not JSON ({e})") from e
|
|
165
|
+
rows = _dig(payload, self.results_path)
|
|
166
|
+
http = HttpBackend("", self.results_path, self.id_field, self.score_field)
|
|
167
|
+
return [http._hit(row, command) for row in rows]
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
@dataclass
|
|
171
|
+
class RunFileBackend(Backend):
|
|
172
|
+
"""A TREC run file: `query_id Q0 doc_id rank score run_name`, whitespace separated.
|
|
173
|
+
|
|
174
|
+
The format the IR community already exchanges results in, so a corpus can be ranked
|
|
175
|
+
once — offline, on a GPU, in whatever language — and evaluated as often as needed with
|
|
176
|
+
nothing running. Queries are addressed by id, so the gold set must supply ids too;
|
|
177
|
+
`--mode` is ignored, since a run file holds exactly one system's output.
|
|
178
|
+
|
|
179
|
+
Lines are sorted by the rank column when present, and otherwise kept in file order.
|
|
180
|
+
"""
|
|
181
|
+
|
|
182
|
+
path: str | Path
|
|
183
|
+
_by_query: dict[str, list[Hit]] = field(init=False, default_factory=dict)
|
|
184
|
+
|
|
185
|
+
def __post_init__(self) -> None:
|
|
186
|
+
try:
|
|
187
|
+
text = Path(self.path).read_text(encoding="utf-8")
|
|
188
|
+
except OSError as e:
|
|
189
|
+
# Surfaced as a backend error so the CLI exits 2 (a usage problem) rather than
|
|
190
|
+
# 1 (a gate breach) — the two mean very different things to a CI job.
|
|
191
|
+
raise BackendError(f"cannot read run file {self.path}: {e}") from e
|
|
192
|
+
staged: dict[str, list[tuple[int, Hit]]] = {}
|
|
193
|
+
for number, line in enumerate(text.splitlines(), 1):
|
|
194
|
+
parts = line.split()
|
|
195
|
+
if not parts or line.lstrip().startswith("#"):
|
|
196
|
+
continue
|
|
197
|
+
if len(parts) < 3:
|
|
198
|
+
raise BackendError(f"{self.path}:{number}: expected at least 3 columns")
|
|
199
|
+
query_id, doc_id = parts[0], parts[2]
|
|
200
|
+
rank = _maybe_int(parts[3]) if len(parts) > 3 else None
|
|
201
|
+
score = _maybe_float(parts[4]) if len(parts) > 4 else None
|
|
202
|
+
staged.setdefault(query_id, []).append(
|
|
203
|
+
(rank if rank is not None else len(staged.get(query_id, [])), Hit(doc_id, score)))
|
|
204
|
+
self._by_query = {q: [hit for _, hit in sorted(rows, key=lambda r: r[0])]
|
|
205
|
+
for q, rows in staged.items()}
|
|
206
|
+
|
|
207
|
+
def search(self, query: str, mode: str) -> list[Hit]:
|
|
208
|
+
if query not in self._by_query:
|
|
209
|
+
raise BackendError(f"{self.path}: no rankings for query id '{query}'. "
|
|
210
|
+
"A run file is addressed by query id — does the gold set "
|
|
211
|
+
"supply ids matching it?")
|
|
212
|
+
return self._by_query[query]
|
|
213
|
+
|
|
214
|
+
@property
|
|
215
|
+
def query_ids(self) -> list[str]:
|
|
216
|
+
return list(self._by_query)
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
def _maybe_int(text: str) -> int | None:
|
|
220
|
+
try:
|
|
221
|
+
return int(text)
|
|
222
|
+
except ValueError:
|
|
223
|
+
return None
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
def _maybe_float(text: str) -> float | None:
|
|
227
|
+
try:
|
|
228
|
+
return float(text)
|
|
229
|
+
except ValueError:
|
|
230
|
+
return None
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
def build_backend(kind: str, **options) -> Backend:
|
|
234
|
+
"""Construct a backend from CLI-shaped options, ignoring the ones it does not take."""
|
|
235
|
+
if kind == "http":
|
|
236
|
+
return HttpBackend(
|
|
237
|
+
url_template=options["url_template"],
|
|
238
|
+
results_path=options.get("results_path", "results"),
|
|
239
|
+
id_field=options.get("id_field", "docId"),
|
|
240
|
+
score_field=options.get("score_field", "score"),
|
|
241
|
+
timeout=options.get("timeout", 120.0),
|
|
242
|
+
)
|
|
243
|
+
if kind == "command":
|
|
244
|
+
return CommandBackend(
|
|
245
|
+
command_template=options["command"],
|
|
246
|
+
results_path=options.get("results_path", "results"),
|
|
247
|
+
id_field=options.get("id_field", "docId"),
|
|
248
|
+
score_field=options.get("score_field", "score"),
|
|
249
|
+
timeout=options.get("timeout", 120.0),
|
|
250
|
+
)
|
|
251
|
+
if kind == "run-file":
|
|
252
|
+
return RunFileBackend(options["run_file"])
|
|
253
|
+
raise BackendError(f"unknown backend '{kind}' — expected http, command or run-file")
|