java-codebase-rag 0.10.2__py3-none-any.whl → 0.11.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.
- java_codebase_rag/cli.py +29 -2
- java_codebase_rag/eval/__init__.py +1 -0
- java_codebase_rag/eval/ground_truth.py +100 -0
- java_codebase_rag/eval/metrics.py +107 -0
- java_codebase_rag/eval/runner.py +556 -0
- java_codebase_rag/search/search_lancedb.py +252 -57
- java_codebase_rag/search/search_lexical.py +10 -0
- java_codebase_rag/search/search_scoring.py +79 -2
- {java_codebase_rag-0.10.2.dist-info → java_codebase_rag-0.11.0.dist-info}/METADATA +1 -1
- {java_codebase_rag-0.10.2.dist-info → java_codebase_rag-0.11.0.dist-info}/RECORD +14 -10
- {java_codebase_rag-0.10.2.dist-info → java_codebase_rag-0.11.0.dist-info}/WHEEL +0 -0
- {java_codebase_rag-0.10.2.dist-info → java_codebase_rag-0.11.0.dist-info}/entry_points.txt +0 -0
- {java_codebase_rag-0.10.2.dist-info → java_codebase_rag-0.11.0.dist-info}/licenses/LICENSE +0 -0
- {java_codebase_rag-0.10.2.dist-info → java_codebase_rag-0.11.0.dist-info}/top_level.txt +0 -0
|
@@ -0,0 +1,556 @@
|
|
|
1
|
+
"""Eval runner — index a corpus, sweep RankConfigs, emit recall/precision/MRR.
|
|
2
|
+
|
|
3
|
+
This is the integration layer of the eval harness (Task 7 of the hybrid-BM25
|
|
4
|
+
RRF plan). Unlike ``eval/metrics.py`` and ``eval/ground_truth.py`` (pure
|
|
5
|
+
stdlib), the runner MAY import the full vector stack (torch / lancedb /
|
|
6
|
+
sentence_transformers) and invokes the operator CLI to build a real index.
|
|
7
|
+
|
|
8
|
+
Pipeline (``run_eval``):
|
|
9
|
+
|
|
10
|
+
1. Build a fresh index into ``cfg.index_dir`` via the operator CLI
|
|
11
|
+
(``java-codebase-rag init``) as a **subprocess** — the stable operator
|
|
12
|
+
surface. Reaching into cocoindex/pipeline internals is fragile (process env,
|
|
13
|
+
progress renderers), so the subprocess wins on robustness. A non-zero exit
|
|
14
|
+
surfaces stdout/stderr in the raised ``RuntimeError``.
|
|
15
|
+
2. Open the index for query: set ``JAVA_CODEBASE_RAG_INDEX_DIR`` so
|
|
16
|
+
``resolve_ladybug_path`` + ``run_search``'s URI resolve to our temp index;
|
|
17
|
+
load ``SentenceTransformer`` once and reuse.
|
|
18
|
+
3. Enumerate ``Symbol`` nodes from the LadybugDB graph (mirrors
|
|
19
|
+
``search_lexical``) and build Tier-A ground truth; optionally concat Tier-B.
|
|
20
|
+
4. For each ``RankConfig`` (``BASELINE_2LIST_CONFIG`` + a 3-list config per
|
|
21
|
+
swept ``k``), run ``run_search`` per query, map rows to ``primary_type_fqn``
|
|
22
|
+
and compute per-query recall@k / precision@k / MRR via ``eval.metrics``.
|
|
23
|
+
5. Aggregate, persist Markdown + JSON under
|
|
24
|
+
``<results_dir>/<ISO-timestamp>/report.{md,json}`` (timestamped subdir so
|
|
25
|
+
successive sweeps don't clobber each other).
|
|
26
|
+
|
|
27
|
+
Granularity note (metric mapping)
|
|
28
|
+
---------------------------------
|
|
29
|
+
Tier-A ``build_tier_a`` sets ``relevant = {symbol.fqn}`` where the fqn may be a
|
|
30
|
+
MEMBER fqn (``com.x.A#processClientMessage()``). ``run_search`` rows carry
|
|
31
|
+
``primary_type_fqn`` = the enclosing TYPE fqn (``com.x.A``, no ``#``). To make
|
|
32
|
+
both sides type-level, the runner normalizes member→type via
|
|
33
|
+
``search_lexical._enclosing_type_fqn`` BEFORE scoring. This keeps Task 6's
|
|
34
|
+
``build_tier_a`` untouched.
|
|
35
|
+
"""
|
|
36
|
+
|
|
37
|
+
from __future__ import annotations
|
|
38
|
+
|
|
39
|
+
import argparse
|
|
40
|
+
import json
|
|
41
|
+
import os
|
|
42
|
+
import subprocess
|
|
43
|
+
import sys
|
|
44
|
+
import tempfile
|
|
45
|
+
import time
|
|
46
|
+
import traceback
|
|
47
|
+
from dataclasses import asdict, dataclass, field
|
|
48
|
+
from datetime import datetime, timezone
|
|
49
|
+
from pathlib import Path
|
|
50
|
+
from typing import Any
|
|
51
|
+
|
|
52
|
+
from java_codebase_rag.eval.ground_truth import LabeledQuery, build_tier_a, load_tier_b
|
|
53
|
+
from java_codebase_rag.eval import metrics as M
|
|
54
|
+
from java_codebase_rag.graph.ladybug_queries import LadybugGraph, resolve_ladybug_path
|
|
55
|
+
from java_codebase_rag.search.index_common import SBERT_MODEL
|
|
56
|
+
from java_codebase_rag.search.search_lexical import _enclosing_type_fqn, _SYMBOL_RETURN
|
|
57
|
+
from java_codebase_rag.search.search_lancedb import run_search
|
|
58
|
+
from java_codebase_rag.search.search_scoring import (
|
|
59
|
+
BASELINE_2LIST_CONFIG,
|
|
60
|
+
RankConfig,
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
# Markdown table columns — order-stable, mirrored by the test suite.
|
|
64
|
+
METRIC_COLUMNS: tuple[str, ...] = (
|
|
65
|
+
"recall@1",
|
|
66
|
+
"recall@5",
|
|
67
|
+
"recall@10",
|
|
68
|
+
"recall@20",
|
|
69
|
+
"precision@5",
|
|
70
|
+
"mrr",
|
|
71
|
+
"p50_latency_ms",
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
# Type-level Symbol kinds emitted by the graph (ast_java._TYPE_KINDS +
|
|
76
|
+
# build_ast_graph._TYPE_KINDS). Confirmed literals — note it's "annotation",
|
|
77
|
+
# NOT "annotation_type". Tier-A recall is measured at type level (chunks carry
|
|
78
|
+
# primary_type_fqn), so member-symbol queries are redundant; defaulting
|
|
79
|
+
# symbol_kinds to these type kinds both bounds fan-out and is semantically right.
|
|
80
|
+
_TYPE_SYMBOL_KINDS: tuple[str, ...] = (
|
|
81
|
+
"class",
|
|
82
|
+
"interface",
|
|
83
|
+
"enum",
|
|
84
|
+
"annotation",
|
|
85
|
+
"record",
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
@dataclass(frozen=True)
|
|
90
|
+
class EvalConfig:
|
|
91
|
+
"""Configuration for a single eval run.
|
|
92
|
+
|
|
93
|
+
``index_dir`` may be empty — the runner creates a temp dir in that case
|
|
94
|
+
(and writes it back into the returned ``EvalReport``).
|
|
95
|
+
"""
|
|
96
|
+
|
|
97
|
+
corpus_dir: str = field(
|
|
98
|
+
default_factory=lambda: str(Path.home() / "jrag-bench" / "shopizer")
|
|
99
|
+
)
|
|
100
|
+
index_dir: str = ""
|
|
101
|
+
results_dir: str = field(
|
|
102
|
+
default_factory=lambda: str(Path.home() / "jrag-bench" / "shopizer" / "results")
|
|
103
|
+
)
|
|
104
|
+
tier_b_path: str | None = None
|
|
105
|
+
ks: tuple[int, ...] = (30, 60, 90, 120)
|
|
106
|
+
top_k_metrics: tuple[int, ...] = (1, 5, 10, 20)
|
|
107
|
+
model_name: str = SBERT_MODEL
|
|
108
|
+
device: str | None = field(
|
|
109
|
+
default_factory=lambda: os.environ.get("SBERT_DEVICE") or None
|
|
110
|
+
)
|
|
111
|
+
# When set, _enumerate_symbols filters Symbol nodes to these kinds (type
|
|
112
|
+
# level by default). None = all kinds (escape hatch). Bounds Tier-A fan-out
|
|
113
|
+
# AND matches the type-level recall granularity.
|
|
114
|
+
symbol_kinds: tuple[str, ...] | None = _TYPE_SYMBOL_KINDS
|
|
115
|
+
# Deterministic cap on Tier-A LabeledQuery items produced (after the kind
|
|
116
|
+
# filter). Tier-B queries are NOT capped (operator-curated). See run_eval.
|
|
117
|
+
max_queries: int = 400
|
|
118
|
+
|
|
119
|
+
def __post_init__(self) -> None:
|
|
120
|
+
if self.max_queries < 1:
|
|
121
|
+
raise ValueError(
|
|
122
|
+
f"EvalConfig.max_queries must be >= 1, got {self.max_queries}"
|
|
123
|
+
)
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
@dataclass(frozen=True)
|
|
127
|
+
class ConfigMetrics:
|
|
128
|
+
"""Aggregated metrics for one RankConfig under test."""
|
|
129
|
+
|
|
130
|
+
config_name: str
|
|
131
|
+
metrics: dict[str, float]
|
|
132
|
+
num_queries: int
|
|
133
|
+
rrf_k: int
|
|
134
|
+
lists: tuple[str, ...]
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
@dataclass(frozen=True)
|
|
138
|
+
class EvalReport:
|
|
139
|
+
"""Result of a full eval sweep."""
|
|
140
|
+
|
|
141
|
+
configs: list[ConfigMetrics]
|
|
142
|
+
timestamp: str
|
|
143
|
+
num_queries: int
|
|
144
|
+
corpus_dir: str
|
|
145
|
+
index_dir: str
|
|
146
|
+
# Tier-A queries available BEFORE the max_queries cap was applied (Tier-B
|
|
147
|
+
# excluded from this count). Recorded so capped runs stay interpretable.
|
|
148
|
+
num_queries_available: int = 0
|
|
149
|
+
# Absolute path to the timestamped output dir holding report.md / report.json.
|
|
150
|
+
out_dir: str = ""
|
|
151
|
+
|
|
152
|
+
def to_json(self) -> str:
|
|
153
|
+
return json.dumps(asdict(self), indent=2, sort_keys=True)
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
# ---------------------------------------------------------------------------
|
|
157
|
+
# Index build (subprocess over the operator CLI — the stable surface)
|
|
158
|
+
# ---------------------------------------------------------------------------
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
def _build_index_subprocess(*, corpus_dir: str, index_dir: str) -> None:
|
|
162
|
+
"""Build a fresh index via ``java-codebase-rag init``.
|
|
163
|
+
|
|
164
|
+
Raises FileNotFoundError if the corpus is missing, RuntimeError on a
|
|
165
|
+
non-zero CLI exit (surfacing clipped stdout/stderr).
|
|
166
|
+
"""
|
|
167
|
+
if not Path(corpus_dir).is_dir():
|
|
168
|
+
raise FileNotFoundError(
|
|
169
|
+
f"eval corpus_dir does not exist: {corpus_dir} "
|
|
170
|
+
"(point EvalConfig.corpus_dir at a checked-out Java repo)"
|
|
171
|
+
)
|
|
172
|
+
|
|
173
|
+
Path(index_dir).mkdir(parents=True, exist_ok=True)
|
|
174
|
+
env = {
|
|
175
|
+
**os.environ,
|
|
176
|
+
"JAVA_CODEBASE_RAG_INDEX_DIR": str(Path(index_dir).resolve()),
|
|
177
|
+
"JAVA_CODEBASE_RAG_SOURCE_ROOT": str(Path(corpus_dir).resolve()),
|
|
178
|
+
}
|
|
179
|
+
cmd = [
|
|
180
|
+
sys.executable,
|
|
181
|
+
"-m",
|
|
182
|
+
"java_codebase_rag.cli",
|
|
183
|
+
"init",
|
|
184
|
+
"--source-root",
|
|
185
|
+
str(Path(corpus_dir).resolve()),
|
|
186
|
+
"--index-dir",
|
|
187
|
+
str(Path(index_dir).resolve()),
|
|
188
|
+
"--quiet",
|
|
189
|
+
]
|
|
190
|
+
proc = subprocess.run(
|
|
191
|
+
cmd,
|
|
192
|
+
capture_output=True,
|
|
193
|
+
text=True,
|
|
194
|
+
env=env,
|
|
195
|
+
timeout=int(os.environ.get("JAVA_CODEBASE_RAG_EVAL_INDEX_TIMEOUT", "1800")),
|
|
196
|
+
)
|
|
197
|
+
if proc.returncode != 0:
|
|
198
|
+
raise RuntimeError(
|
|
199
|
+
f"java-codebase-rag init exited {proc.returncode} for corpus {corpus_dir}.\n"
|
|
200
|
+
f"--- stdout (clipped 8000) ---\n{proc.stdout[-8000:]}\n"
|
|
201
|
+
f"--- stderr (clipped 8000) ---\n{proc.stderr[-8000:]}"
|
|
202
|
+
)
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
# ---------------------------------------------------------------------------
|
|
206
|
+
# Symbol enumeration (mirror search_lexical)
|
|
207
|
+
# ---------------------------------------------------------------------------
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
class _SymbolRow:
|
|
211
|
+
"""Attribute-view over a Cypher row dict (build_tier_a ducks on .fqn/.name)."""
|
|
212
|
+
|
|
213
|
+
__slots__ = ("fqn", "name", "kind")
|
|
214
|
+
|
|
215
|
+
def __init__(self, row: dict[str, Any]) -> None:
|
|
216
|
+
self.fqn = str(row.get("fqn") or "")
|
|
217
|
+
self.name = str(row.get("name") or "")
|
|
218
|
+
self.kind = str(row.get("kind") or "")
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
def _enumerate_symbols(
|
|
222
|
+
graph: LadybugGraph, *, symbol_kinds: tuple[str, ...] | None
|
|
223
|
+
) -> list[_SymbolRow]:
|
|
224
|
+
"""Return Symbol rows as duck-typed objects (.fqn / .name) for build_tier_a.
|
|
225
|
+
|
|
226
|
+
``LadybugGraph._rows`` returns dicts; ``build_tier_a``'s ``SymbolLike``
|
|
227
|
+
protocol ducks on attributes, so we wrap each row.
|
|
228
|
+
|
|
229
|
+
When ``symbol_kinds`` is not None, filters to those kinds via a
|
|
230
|
+
parameterized ``WHERE s.kind IN $kinds`` predicate — type-level only by
|
|
231
|
+
default, which both bounds fan-out and matches the type-level recall
|
|
232
|
+
granularity (chunks carry ``primary_type_fqn``).
|
|
233
|
+
"""
|
|
234
|
+
if symbol_kinds is None:
|
|
235
|
+
cypher = f"MATCH (s:Symbol) RETURN {_SYMBOL_RETURN}"
|
|
236
|
+
rows = graph._rows(cypher) # noqa: SLF001
|
|
237
|
+
else:
|
|
238
|
+
cypher = f"MATCH (s:Symbol) WHERE s.kind IN $kinds RETURN {_SYMBOL_RETURN}"
|
|
239
|
+
rows = graph._rows(cypher, {"kinds": list(symbol_kinds)}) # noqa: SLF001
|
|
240
|
+
return [_SymbolRow(r) for r in rows]
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
# ---------------------------------------------------------------------------
|
|
244
|
+
# Metric computation
|
|
245
|
+
# ---------------------------------------------------------------------------
|
|
246
|
+
|
|
247
|
+
|
|
248
|
+
def _retrieved_fqns(rows: list[dict]) -> list[str]:
|
|
249
|
+
"""Map search rows to ordered, deduped type-level FQNs."""
|
|
250
|
+
out: list[str] = []
|
|
251
|
+
seen: set[str] = set()
|
|
252
|
+
for r in rows:
|
|
253
|
+
fqn = r.get("primary_type_fqn") or r.get("fqn") or ""
|
|
254
|
+
if not fqn:
|
|
255
|
+
continue
|
|
256
|
+
if fqn in seen:
|
|
257
|
+
continue
|
|
258
|
+
seen.add(fqn)
|
|
259
|
+
out.append(fqn)
|
|
260
|
+
return out
|
|
261
|
+
|
|
262
|
+
|
|
263
|
+
def _relevant_type_fqns(labeled: LabeledQuery) -> set[str]:
|
|
264
|
+
"""Normalize member FQNs to enclosing-type FQNs so both sides are type-level."""
|
|
265
|
+
return {_enclosing_type_fqn(fqn) for fqn in labeled.relevant if fqn}
|
|
266
|
+
|
|
267
|
+
|
|
268
|
+
def _p50(values: list[float]) -> float:
|
|
269
|
+
if not values:
|
|
270
|
+
return 0.0
|
|
271
|
+
s = sorted(values)
|
|
272
|
+
mid = len(s) // 2
|
|
273
|
+
if len(s) % 2:
|
|
274
|
+
return float(s[mid])
|
|
275
|
+
return float((s[mid - 1] + s[mid]) / 2.0)
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
def _eval_one_config(
|
|
279
|
+
*,
|
|
280
|
+
config_name: str,
|
|
281
|
+
rank_config: RankConfig,
|
|
282
|
+
queries: list[LabeledQuery],
|
|
283
|
+
uri: str,
|
|
284
|
+
ladybug_path: str,
|
|
285
|
+
model,
|
|
286
|
+
model_name: str,
|
|
287
|
+
device: str | None,
|
|
288
|
+
top_k_metrics: tuple[int, ...],
|
|
289
|
+
limit: int,
|
|
290
|
+
) -> ConfigMetrics:
|
|
291
|
+
"""Run one RankConfig over all queries; return aggregated ConfigMetrics."""
|
|
292
|
+
per_query: list[dict[str, float]] = []
|
|
293
|
+
latencies_ms: list[float] = []
|
|
294
|
+
|
|
295
|
+
for q in queries:
|
|
296
|
+
t0 = time.perf_counter()
|
|
297
|
+
rows = run_search(
|
|
298
|
+
q.query,
|
|
299
|
+
uri=uri,
|
|
300
|
+
table_keys=["java"],
|
|
301
|
+
limit=limit,
|
|
302
|
+
offset=0,
|
|
303
|
+
path_substring=None,
|
|
304
|
+
model_name=model_name,
|
|
305
|
+
device=device,
|
|
306
|
+
model=model,
|
|
307
|
+
rank_config=rank_config,
|
|
308
|
+
graph_expand=True,
|
|
309
|
+
expand_depth=1,
|
|
310
|
+
ladybug_path=ladybug_path,
|
|
311
|
+
dedup_by_fqn=True,
|
|
312
|
+
)
|
|
313
|
+
elapsed_ms = (time.perf_counter() - t0) * 1000.0
|
|
314
|
+
latencies_ms.append(elapsed_ms)
|
|
315
|
+
|
|
316
|
+
retrieved = _retrieved_fqns(rows)
|
|
317
|
+
relevant = _relevant_type_fqns(q)
|
|
318
|
+
if not relevant:
|
|
319
|
+
# No relevant set (e.g. Tier-B with empty relevant) — skip from scoring.
|
|
320
|
+
continue
|
|
321
|
+
|
|
322
|
+
qm: dict[str, float] = {}
|
|
323
|
+
for k in top_k_metrics:
|
|
324
|
+
qm[f"recall@{k}"] = M.recall_at_k(retrieved, relevant, k)
|
|
325
|
+
qm["precision@5"] = M.precision_at_k(retrieved, relevant, 5)
|
|
326
|
+
qm["mrr"] = M.reciprocal_rank(retrieved, relevant)
|
|
327
|
+
per_query.append(qm)
|
|
328
|
+
|
|
329
|
+
agg = M.aggregate(per_query)
|
|
330
|
+
metrics: dict[str, float] = {}
|
|
331
|
+
for k in top_k_metrics:
|
|
332
|
+
metrics[f"recall@{k}"] = float(agg.get(f"recall@{k}", 0.0))
|
|
333
|
+
metrics["precision@5"] = float(agg.get("precision@5", 0.0))
|
|
334
|
+
metrics["mrr"] = float(agg.get("mrr", 0.0))
|
|
335
|
+
metrics["p50_latency_ms"] = _p50(latencies_ms)
|
|
336
|
+
|
|
337
|
+
return ConfigMetrics(
|
|
338
|
+
config_name=config_name,
|
|
339
|
+
metrics=metrics,
|
|
340
|
+
num_queries=len(per_query),
|
|
341
|
+
rrf_k=rank_config.rrf_k,
|
|
342
|
+
lists=tuple(sorted(rank_config.lists)),
|
|
343
|
+
)
|
|
344
|
+
|
|
345
|
+
|
|
346
|
+
# ---------------------------------------------------------------------------
|
|
347
|
+
# Markdown rendering
|
|
348
|
+
# ---------------------------------------------------------------------------
|
|
349
|
+
|
|
350
|
+
|
|
351
|
+
def _render_markdown(report: EvalReport) -> str:
|
|
352
|
+
lines: list[str] = []
|
|
353
|
+
lines.append(f"# Eval Report — {report.timestamp}")
|
|
354
|
+
lines.append("")
|
|
355
|
+
lines.append(
|
|
356
|
+
f"Corpus: `{report.corpus_dir}` | Index: `{report.index_dir}` "
|
|
357
|
+
f"| Queries scored: {report.num_queries}"
|
|
358
|
+
)
|
|
359
|
+
if report.num_queries_available:
|
|
360
|
+
lines.append(
|
|
361
|
+
f"Tier-A available (pre-cap): {report.num_queries_available} | "
|
|
362
|
+
f"Total scored (Tier-A kept + Tier-B): {report.num_queries}"
|
|
363
|
+
)
|
|
364
|
+
lines.append("")
|
|
365
|
+
header = "| config | " + " | ".join(METRIC_COLUMNS) + " |"
|
|
366
|
+
sep = "| --- " * (len(METRIC_COLUMNS) + 1) + "|"
|
|
367
|
+
lines.append(header)
|
|
368
|
+
lines.append(sep)
|
|
369
|
+
for entry in report.configs:
|
|
370
|
+
cells = [entry.config_name]
|
|
371
|
+
for col in METRIC_COLUMNS:
|
|
372
|
+
v = entry.metrics.get(col, 0.0)
|
|
373
|
+
if col == "p50_latency_ms":
|
|
374
|
+
cells.append(f"{v:.1f}")
|
|
375
|
+
else:
|
|
376
|
+
cells.append(f"{v:.4f}")
|
|
377
|
+
lines.append("| " + " | ".join(cells) + " |")
|
|
378
|
+
lines.append("")
|
|
379
|
+
return "\n".join(lines)
|
|
380
|
+
|
|
381
|
+
|
|
382
|
+
# ---------------------------------------------------------------------------
|
|
383
|
+
# Orchestration
|
|
384
|
+
# ---------------------------------------------------------------------------
|
|
385
|
+
|
|
386
|
+
|
|
387
|
+
def run_eval(cfg: EvalConfig) -> EvalReport:
|
|
388
|
+
"""Build a fresh index, sweep RankConfigs, return an EvalReport.
|
|
389
|
+
|
|
390
|
+
See module docstring for the pipeline and the metric-granularity note.
|
|
391
|
+
"""
|
|
392
|
+
# Late import — torch/lancedb only needed for the run, not for module import.
|
|
393
|
+
from sentence_transformers import SentenceTransformer
|
|
394
|
+
|
|
395
|
+
# Resolve index_dir (temp dir if blank).
|
|
396
|
+
if not cfg.index_dir:
|
|
397
|
+
index_dir = tempfile.mkdtemp(prefix="jrag-eval-")
|
|
398
|
+
else:
|
|
399
|
+
index_dir = cfg.index_dir
|
|
400
|
+
Path(index_dir).mkdir(parents=True, exist_ok=True)
|
|
401
|
+
|
|
402
|
+
# 1. Build the index (subprocess).
|
|
403
|
+
_build_index_subprocess(corpus_dir=cfg.corpus_dir, index_dir=index_dir)
|
|
404
|
+
|
|
405
|
+
# 2. Wire the process env so resolve_ladybug_path + run_search's URI hit our index.
|
|
406
|
+
os.environ["JAVA_CODEBASE_RAG_INDEX_DIR"] = str(Path(index_dir).resolve())
|
|
407
|
+
os.environ.setdefault(
|
|
408
|
+
"JAVA_CODEBASE_RAG_SOURCE_ROOT", str(Path(cfg.corpus_dir).resolve())
|
|
409
|
+
)
|
|
410
|
+
uri = str(Path(index_dir).resolve())
|
|
411
|
+
ladybug_path = resolve_ladybug_path(None)
|
|
412
|
+
|
|
413
|
+
# Load the model ONCE — pass into every run_search call.
|
|
414
|
+
model = SentenceTransformer(
|
|
415
|
+
cfg.model_name, device=cfg.device, trust_remote_code=True
|
|
416
|
+
)
|
|
417
|
+
|
|
418
|
+
# 3. Open graph + build ground truth.
|
|
419
|
+
# Reset the LadybugGraph singleton — a prior test/process may have cached
|
|
420
|
+
# a different path. We force-bind to our index's graph.
|
|
421
|
+
LadybugGraph.reset_for_path(None)
|
|
422
|
+
graph = LadybugGraph.get(ladybug_path)
|
|
423
|
+
|
|
424
|
+
symbols = _enumerate_symbols(graph, symbol_kinds=cfg.symbol_kinds)
|
|
425
|
+
tier_a = list(build_tier_a(symbols))
|
|
426
|
+
num_queries_available = len(tier_a)
|
|
427
|
+
# Deterministic cap on Tier-A: sort by (query, fqn) and keep the first
|
|
428
|
+
# max_queries. fqn lives as the sole element of each LabeledQuery.relevant
|
|
429
|
+
# (build_tier_a sets relevant = {symbol.fqn}). Tier-B is operator-curated
|
|
430
|
+
# and NOT capped.
|
|
431
|
+
tier_a_sorted = sorted(tier_a, key=lambda q: (q.query, next(iter(q.relevant))))
|
|
432
|
+
tier_a_kept = tier_a_sorted[: cfg.max_queries]
|
|
433
|
+
queries = list(tier_a_kept)
|
|
434
|
+
# Tier-B is optional: a configured but missing path means "Tier-B disabled"
|
|
435
|
+
# (matches load_tier_b's docstring). load_tier_b itself still raises
|
|
436
|
+
# FileNotFoundError when called directly on a genuinely missing path.
|
|
437
|
+
if cfg.tier_b_path and Path(cfg.tier_b_path).exists():
|
|
438
|
+
queries.extend(load_tier_b(cfg.tier_b_path))
|
|
439
|
+
|
|
440
|
+
# 4. Enumerate configs: BASELINE_2LIST_CONFIG (k=60) + 3-list at each swept k.
|
|
441
|
+
limit = max(cfg.top_k_metrics)
|
|
442
|
+
configs: list[tuple[str, RankConfig]] = [
|
|
443
|
+
("baseline_2list_k60", BASELINE_2LIST_CONFIG),
|
|
444
|
+
]
|
|
445
|
+
for k in cfg.ks:
|
|
446
|
+
configs.append(
|
|
447
|
+
(
|
|
448
|
+
f"hybrid_3list_k{k}",
|
|
449
|
+
RankConfig(
|
|
450
|
+
lists=frozenset({"vector", "graph", "bm25"}),
|
|
451
|
+
rrf_k=k,
|
|
452
|
+
),
|
|
453
|
+
)
|
|
454
|
+
)
|
|
455
|
+
|
|
456
|
+
# 5. Run + aggregate.
|
|
457
|
+
results: list[ConfigMetrics] = []
|
|
458
|
+
for name, rc in configs:
|
|
459
|
+
results.append(
|
|
460
|
+
_eval_one_config(
|
|
461
|
+
config_name=name,
|
|
462
|
+
rank_config=rc,
|
|
463
|
+
queries=queries,
|
|
464
|
+
uri=uri,
|
|
465
|
+
ladybug_path=ladybug_path,
|
|
466
|
+
model=model,
|
|
467
|
+
model_name=cfg.model_name,
|
|
468
|
+
device=cfg.device,
|
|
469
|
+
top_k_metrics=cfg.top_k_metrics,
|
|
470
|
+
limit=limit,
|
|
471
|
+
)
|
|
472
|
+
)
|
|
473
|
+
|
|
474
|
+
timestamp = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
|
|
475
|
+
# Namespace outputs under the timestamp so successive sweeps don't clobber.
|
|
476
|
+
out_dir = str(Path(cfg.results_dir) / timestamp)
|
|
477
|
+
report = EvalReport(
|
|
478
|
+
configs=results,
|
|
479
|
+
timestamp=timestamp,
|
|
480
|
+
num_queries=len(queries),
|
|
481
|
+
corpus_dir=cfg.corpus_dir,
|
|
482
|
+
index_dir=index_dir,
|
|
483
|
+
num_queries_available=num_queries_available,
|
|
484
|
+
out_dir=out_dir,
|
|
485
|
+
)
|
|
486
|
+
|
|
487
|
+
# 6. Persist into <results_dir>/<timestamp>/report.{md,json}.
|
|
488
|
+
_persist(report, cfg.results_dir)
|
|
489
|
+
return report
|
|
490
|
+
|
|
491
|
+
|
|
492
|
+
def _persist(report: EvalReport, results_dir: str) -> None:
|
|
493
|
+
# Write into a timestamped subdir so successive sweeps don't clobber.
|
|
494
|
+
out = Path(results_dir) / report.timestamp
|
|
495
|
+
out.mkdir(parents=True, exist_ok=True)
|
|
496
|
+
(out / "report.md").write_text(_render_markdown(report))
|
|
497
|
+
(out / "report.json").write_text(report.to_json())
|
|
498
|
+
|
|
499
|
+
|
|
500
|
+
def _build_eval_config_from_args(argv: list[str] | None = None) -> EvalConfig:
|
|
501
|
+
"""Parse CLI args into an EvalConfig. Defaults mirror EvalConfig fields."""
|
|
502
|
+
# Start from defaults so unspecified flags inherit EvalConfig's own defaults.
|
|
503
|
+
base = EvalConfig()
|
|
504
|
+
parser = argparse.ArgumentParser(
|
|
505
|
+
prog="python -m java_codebase_rag.eval.runner",
|
|
506
|
+
description="Index a corpus, sweep RankConfigs, emit recall/precision/MRR.",
|
|
507
|
+
)
|
|
508
|
+
parser.add_argument(
|
|
509
|
+
"corpus_dir",
|
|
510
|
+
nargs="?",
|
|
511
|
+
default=base.corpus_dir,
|
|
512
|
+
help=f"Java repo to index (default: {base.corpus_dir})",
|
|
513
|
+
)
|
|
514
|
+
parser.add_argument("--index-dir", default=base.index_dir,
|
|
515
|
+
help="Lance+Ladybug index dir (default: temp dir).")
|
|
516
|
+
parser.add_argument("--results-dir", default=base.results_dir,
|
|
517
|
+
help=f"Where to write report.{{md,json}} (default: {base.results_dir})")
|
|
518
|
+
parser.add_argument("--max-queries", type=int, default=base.max_queries,
|
|
519
|
+
help=f"Cap on Tier-A queries (default: {base.max_queries}).")
|
|
520
|
+
parser.add_argument(
|
|
521
|
+
"--ks", default=",".join(str(k) for k in base.ks),
|
|
522
|
+
help="Comma-separated RRF k constants to sweep (default: %(default)s).",
|
|
523
|
+
)
|
|
524
|
+
parser.add_argument("--tier-b", default=base.tier_b_path,
|
|
525
|
+
help="Optional path to a Tier-B ground-truth file (missing ⇒ disabled).")
|
|
526
|
+
parser.add_argument("--device", default=base.device,
|
|
527
|
+
help="SBERT device (default: SBERT_DEVICE env or auto).")
|
|
528
|
+
args = parser.parse_args(argv)
|
|
529
|
+
|
|
530
|
+
try:
|
|
531
|
+
ks = tuple(int(k.strip()) for k in args.ks.split(",") if k.strip())
|
|
532
|
+
except ValueError:
|
|
533
|
+
parser.error(f"--ks must be comma-separated ints, got {args.ks!r}")
|
|
534
|
+
if not ks:
|
|
535
|
+
parser.error("--ks must contain at least one value")
|
|
536
|
+
|
|
537
|
+
return EvalConfig(
|
|
538
|
+
corpus_dir=args.corpus_dir,
|
|
539
|
+
index_dir=args.index_dir,
|
|
540
|
+
results_dir=args.results_dir,
|
|
541
|
+
tier_b_path=args.tier_b,
|
|
542
|
+
ks=ks,
|
|
543
|
+
max_queries=args.max_queries,
|
|
544
|
+
device=args.device,
|
|
545
|
+
)
|
|
546
|
+
|
|
547
|
+
|
|
548
|
+
if __name__ == "__main__":
|
|
549
|
+
cfg = _build_eval_config_from_args()
|
|
550
|
+
try:
|
|
551
|
+
report = run_eval(cfg)
|
|
552
|
+
except Exception:
|
|
553
|
+
traceback.print_exc()
|
|
554
|
+
sys.exit(1)
|
|
555
|
+
print(report.out_dir)
|
|
556
|
+
sys.exit(0)
|