pretensor 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.
- pretensor/__init__.py +50 -0
- pretensor/benchmark/__init__.py +54 -0
- pretensor/benchmark/cli.py +294 -0
- pretensor/benchmark/fixtures.py +84 -0
- pretensor/benchmark/l1/__init__.py +23 -0
- pretensor/benchmark/l1/metrics.py +141 -0
- pretensor/benchmark/l1/pipeline.py +188 -0
- pretensor/benchmark/l1/runner.py +245 -0
- pretensor/benchmark/l2/__init__.py +27 -0
- pretensor/benchmark/l2/gold.py +236 -0
- pretensor/benchmark/l2/metrics.py +146 -0
- pretensor/benchmark/l2/pipeline.py +124 -0
- pretensor/benchmark/l2/runner.py +530 -0
- pretensor/benchmark/l3/__init__.py +73 -0
- pretensor/benchmark/l3/agent.py +316 -0
- pretensor/benchmark/l3/db.py +188 -0
- pretensor/benchmark/l3/gold.py +85 -0
- pretensor/benchmark/l3/llm_client.py +395 -0
- pretensor/benchmark/l3/mcp_client.py +357 -0
- pretensor/benchmark/l3/pretensor_runner.py +456 -0
- pretensor/benchmark/l3/prompt.py +132 -0
- pretensor/benchmark/l3/runner.py +358 -0
- pretensor/benchmark/l3/sql_equivalence.py +176 -0
- pretensor/benchmark/release_gate.py +448 -0
- pretensor/benchmark/results.py +298 -0
- pretensor/benchmark/runner.py +109 -0
- pretensor/cli/__init__.py +1 -0
- pretensor/cli/commands/_source_runner.py +147 -0
- pretensor/cli/commands/analyze.py +201 -0
- pretensor/cli/commands/connections/__init__.py +7 -0
- pretensor/cli/commands/connections/add_remove.py +126 -0
- pretensor/cli/commands/connections/register.py +12 -0
- pretensor/cli/commands/export.py +131 -0
- pretensor/cli/commands/index.py +559 -0
- pretensor/cli/commands/list.py +76 -0
- pretensor/cli/commands/quickstart.py +207 -0
- pretensor/cli/commands/reindex.py +646 -0
- pretensor/cli/commands/semantic.py +190 -0
- pretensor/cli/commands/serve.py +144 -0
- pretensor/cli/commands/sync_grants.py +149 -0
- pretensor/cli/commands/validate.py +176 -0
- pretensor/cli/config_file.py +442 -0
- pretensor/cli/constants.py +10 -0
- pretensor/cli/dbt_enrichment.py +96 -0
- pretensor/cli/main.py +109 -0
- pretensor/cli/paths.py +43 -0
- pretensor/cli/plugin.py +52 -0
- pretensor/config.py +226 -0
- pretensor/connectors/__init__.py +29 -0
- pretensor/connectors/base.py +165 -0
- pretensor/connectors/bigquery.py +468 -0
- pretensor/connectors/inspect.py +321 -0
- pretensor/connectors/lineage_sqlglot.py +97 -0
- pretensor/connectors/models.py +130 -0
- pretensor/connectors/mysql.py +402 -0
- pretensor/connectors/pg_array_parse.py +53 -0
- pretensor/connectors/postgres.py +938 -0
- pretensor/connectors/registry.py +93 -0
- pretensor/connectors/snapshot.py +244 -0
- pretensor/connectors/snowflake.py +908 -0
- pretensor/core/__init__.py +1 -0
- pretensor/core/builder.py +307 -0
- pretensor/core/dsn_crypto.py +51 -0
- pretensor/core/graph_schema_manager.py +246 -0
- pretensor/core/graph_store.py +1226 -0
- pretensor/core/ids.py +101 -0
- pretensor/core/portable_export.py +276 -0
- pretensor/core/query_runner.py +67 -0
- pretensor/core/registry.py +209 -0
- pretensor/core/schema.py +473 -0
- pretensor/core/secure_io.py +93 -0
- pretensor/core/store.py +469 -0
- pretensor/enrichment/__init__.py +1 -0
- pretensor/enrichment/analyze/__init__.py +0 -0
- pretensor/enrichment/analyze/classify.py +49 -0
- pretensor/enrichment/analyze/extract_python.py +196 -0
- pretensor/enrichment/analyze/parse.py +141 -0
- pretensor/enrichment/analyze/pipeline.py +195 -0
- pretensor/enrichment/analyze/summary.py +38 -0
- pretensor/enrichment/analyze/walker.py +98 -0
- pretensor/enrichment/analyze/writers.py +214 -0
- pretensor/enrichment/dbt/__init__.py +30 -0
- pretensor/enrichment/dbt/lineage.py +100 -0
- pretensor/enrichment/dbt/manifest.py +300 -0
- pretensor/enrichment/dbt/metadata.py +263 -0
- pretensor/enrichment/dbt/pipeline.py +77 -0
- pretensor/enrichment/dbt/resolution.py +101 -0
- pretensor/enrichment/dbt/signals.py +305 -0
- pretensor/entities/__init__.py +27 -0
- pretensor/entities/builder.py +63 -0
- pretensor/entities/classifier.py +383 -0
- pretensor/entities/llm_extract.py +66 -0
- pretensor/errors.py +35 -0
- pretensor/graph_models/__init__.py +17 -0
- pretensor/graph_models/base.py +11 -0
- pretensor/graph_models/consumer.py +71 -0
- pretensor/graph_models/edge.py +35 -0
- pretensor/graph_models/entity.py +21 -0
- pretensor/graph_models/node.py +79 -0
- pretensor/graph_models/relationship.py +33 -0
- pretensor/integrations/__init__.py +42 -0
- pretensor/integrations/_base.py +138 -0
- pretensor/integrations/google_adk.py +49 -0
- pretensor/integrations/langchain.py +55 -0
- pretensor/integrations/llamaindex.py +53 -0
- pretensor/intelligence/__init__.py +33 -0
- pretensor/intelligence/cluster_labeler.py +425 -0
- pretensor/intelligence/clustering.py +168 -0
- pretensor/intelligence/combining.py +32 -0
- pretensor/intelligence/discovery.py +114 -0
- pretensor/intelligence/embeddings.py +317 -0
- pretensor/intelligence/graph_export.py +200 -0
- pretensor/intelligence/heuristic.py +544 -0
- pretensor/intelligence/join_paths/__init__.py +130 -0
- pretensor/intelligence/join_paths/on_demand.py +516 -0
- pretensor/intelligence/join_paths/storage.py +70 -0
- pretensor/intelligence/llm_infer.py +78 -0
- pretensor/intelligence/llm_runtime.py +62 -0
- pretensor/intelligence/metric_templates.py +193 -0
- pretensor/intelligence/pipeline.py +364 -0
- pretensor/intelligence/role_exemplars.py +263 -0
- pretensor/intelligence/schema_classification.py +360 -0
- pretensor/intelligence/scoring.py +76 -0
- pretensor/intelligence/semantic.py +240 -0
- pretensor/intelligence/shadow_alias.py +101 -0
- pretensor/intelligence/statistical.py +50 -0
- pretensor/intelligence/steps.py +191 -0
- pretensor/intelligence/steps_embedding.py +168 -0
- pretensor/introspection/__init__.py +6 -0
- pretensor/introspection/inspector.py +5 -0
- pretensor/introspection/models/__init__.py +0 -0
- pretensor/introspection/models/base.py +5 -0
- pretensor/introspection/models/config.py +237 -0
- pretensor/introspection/models/dsn.py +550 -0
- pretensor/introspection/models/plan.py +116 -0
- pretensor/introspection/models/schema.py +10 -0
- pretensor/introspection/models/semantic.py +121 -0
- pretensor/introspection/models/validation.py +116 -0
- pretensor/introspection/snapshot.py +46 -0
- pretensor/mcp/__init__.py +16 -0
- pretensor/mcp/config_json.py +24 -0
- pretensor/mcp/payload_types.py +274 -0
- pretensor/mcp/resources/__init__.py +17 -0
- pretensor/mcp/resources/markdown.py +314 -0
- pretensor/mcp/server.py +285 -0
- pretensor/mcp/service.py +49 -0
- pretensor/mcp/service_context.py +142 -0
- pretensor/mcp/service_registry.py +294 -0
- pretensor/mcp/store_cache.py +43 -0
- pretensor/mcp/tool_registry.py +136 -0
- pretensor/mcp/tools/__init__.py +1 -0
- pretensor/mcp/tools/_rank.py +244 -0
- pretensor/mcp/tools/_timed.py +26 -0
- pretensor/mcp/tools/compile_metric.py +144 -0
- pretensor/mcp/tools/consumers.py +161 -0
- pretensor/mcp/tools/context.py +1121 -0
- pretensor/mcp/tools/cypher.py +509 -0
- pretensor/mcp/tools/detect_changes.py +254 -0
- pretensor/mcp/tools/impact.py +271 -0
- pretensor/mcp/tools/list.py +131 -0
- pretensor/mcp/tools/schema.py +170 -0
- pretensor/mcp/tools/search.py +316 -0
- pretensor/mcp/tools/semantic_search.py +282 -0
- pretensor/mcp/tools/traverse.py +1027 -0
- pretensor/mcp/tools/validate_sql.py +150 -0
- pretensor/observability.py +203 -0
- pretensor/py.typed +0 -0
- pretensor/quickstart/README.md +29 -0
- pretensor/quickstart/__init__.py +6 -0
- pretensor/quickstart/docker-compose.yml +18 -0
- pretensor/quickstart/pagila_data.sql +63 -0
- pretensor/quickstart/pagila_ddl.sql +92 -0
- pretensor/search/__init__.py +6 -0
- pretensor/search/base.py +80 -0
- pretensor/search/index.py +435 -0
- pretensor/semantic/__init__.py +24 -0
- pretensor/semantic/base.py +123 -0
- pretensor/semantic/compiler.py +487 -0
- pretensor/semantic/yaml_layer.py +180 -0
- pretensor/skills/__init__.py +5 -0
- pretensor/skills/generator.py +235 -0
- pretensor/staleness/__init__.py +15 -0
- pretensor/staleness/graph_patcher.py +355 -0
- pretensor/staleness/impact_analyzer.py +162 -0
- pretensor/staleness/snapshot_store.py +38 -0
- pretensor/validation/__init__.py +9 -0
- pretensor/validation/query_validator.py +436 -0
- pretensor/visibility/__init__.py +23 -0
- pretensor/visibility/config.py +126 -0
- pretensor/visibility/filter.py +143 -0
- pretensor/visibility/kuzu_helpers.py +32 -0
- pretensor/visibility/runtime.py +36 -0
- pretensor/visibility/sync_grants.py +188 -0
- pretensor-0.1.0.dist-info/METADATA +251 -0
- pretensor-0.1.0.dist-info/RECORD +198 -0
- pretensor-0.1.0.dist-info/WHEEL +4 -0
- pretensor-0.1.0.dist-info/entry_points.txt +2 -0
- pretensor-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,530 @@
|
|
|
1
|
+
"""``run_l2`` — collect L2 metrics into a deterministic JSON document.
|
|
2
|
+
|
|
3
|
+
The runner orchestrates the four L2 metrics by:
|
|
4
|
+
|
|
5
|
+
1. Building a fresh in-memory Kuzu graph + ``registry.json`` from the
|
|
6
|
+
fixture's schema YAML (see :mod:`pretensor.benchmark.l2.pipeline`).
|
|
7
|
+
2. Calling the MCP tool payload functions in-process for each gold
|
|
8
|
+
observation — no MCP transport, no subprocess, fully deterministic.
|
|
9
|
+
3. Handing already-collected observations to the pure metric helpers
|
|
10
|
+
in :mod:`pretensor.benchmark.l2.metrics`.
|
|
11
|
+
4. Wrapping the scalar values in :class:`Metric` and emitting a
|
|
12
|
+
:class:`BenchmarkResult` with byte-identical JSON across re-runs.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
import hashlib
|
|
18
|
+
import importlib
|
|
19
|
+
import importlib.metadata
|
|
20
|
+
import json
|
|
21
|
+
import sys
|
|
22
|
+
import tempfile
|
|
23
|
+
from pathlib import Path
|
|
24
|
+
from typing import TYPE_CHECKING, Any
|
|
25
|
+
|
|
26
|
+
from pretensor.benchmark.fixtures import load_dataset
|
|
27
|
+
from pretensor.benchmark.l2.gold import (
|
|
28
|
+
MetricTemplateEntry,
|
|
29
|
+
QueryGoldEntry,
|
|
30
|
+
TraverseGoldEntry,
|
|
31
|
+
bare_table_name,
|
|
32
|
+
load_metric_templates,
|
|
33
|
+
load_query_gold,
|
|
34
|
+
load_traverse_gold,
|
|
35
|
+
)
|
|
36
|
+
from pretensor.benchmark.l2.metrics import (
|
|
37
|
+
JoinPair,
|
|
38
|
+
RankedHit,
|
|
39
|
+
compile_metric_correctness,
|
|
40
|
+
query_recall_at_k,
|
|
41
|
+
semantic_search_recall_at_k,
|
|
42
|
+
top_k_with_ties,
|
|
43
|
+
traverse_correctness,
|
|
44
|
+
)
|
|
45
|
+
from pretensor.benchmark.l2.pipeline import build_l2_graph_dir
|
|
46
|
+
from pretensor.benchmark.results import BenchmarkResult, Metric, write_json
|
|
47
|
+
from pretensor.connectors.models import SchemaSnapshot
|
|
48
|
+
from pretensor.mcp.tools.compile_metric import compile_metric_payload
|
|
49
|
+
from pretensor.mcp.tools.search import query_payload
|
|
50
|
+
from pretensor.mcp.tools.traverse import traverse_payload
|
|
51
|
+
|
|
52
|
+
if TYPE_CHECKING:
|
|
53
|
+
from pretensor.benchmark.runner import Dataset
|
|
54
|
+
|
|
55
|
+
__all__ = ["run_l2"]
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
_DETERMINISTIC_RAN_AT = "1970-01-01T00:00:00Z"
|
|
59
|
+
"""Pinned timestamp so two consecutive runs produce byte-identical JSON."""
|
|
60
|
+
|
|
61
|
+
_RECALL_K = 5
|
|
62
|
+
"""Top-K cutoff used for both ``query`` and ``semantic_search`` Recall@K.
|
|
63
|
+
|
|
64
|
+
The metric key in the JSON output (``query_recall_at_5`` /
|
|
65
|
+
``semantic_search_recall_at_5``) is built from this value.
|
|
66
|
+
"""
|
|
67
|
+
|
|
68
|
+
_EMBEDDINGS_SENTINEL = "onnxruntime"
|
|
69
|
+
"""Importable module that ships with the ``[embeddings]`` extra.
|
|
70
|
+
|
|
71
|
+
Same probe L1 uses; mirrors the ``onnxruntime`` runtime check the
|
|
72
|
+
embedding model relies on.
|
|
73
|
+
"""
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def run_l2(
|
|
77
|
+
dataset: Dataset,
|
|
78
|
+
out: Path | None,
|
|
79
|
+
graph_dir: Path, # noqa: ARG001 — L2 builds a fresh in-process graph
|
|
80
|
+
*,
|
|
81
|
+
embeddings: bool,
|
|
82
|
+
) -> None:
|
|
83
|
+
"""Run the four L2 metrics for ``dataset`` and emit a ``BenchmarkResult``.
|
|
84
|
+
|
|
85
|
+
``graph_dir`` is part of the CLI contract for parity with L1 / L3 but
|
|
86
|
+
unused here — L2 indexes the fixture into a fresh per-call temporary
|
|
87
|
+
store so the output is reproducible from a clean checkout.
|
|
88
|
+
"""
|
|
89
|
+
fixture = load_dataset(dataset)
|
|
90
|
+
snapshot_text = fixture.schema_yaml_path.read_text(encoding="utf-8")
|
|
91
|
+
snapshot = SchemaSnapshot.from_yaml(snapshot_text)
|
|
92
|
+
fixture_sha = (
|
|
93
|
+
"sha256:" + hashlib.sha256(fixture.schema_yaml_path.read_bytes()).hexdigest()
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
notes: list[str] = []
|
|
97
|
+
embeddings_enabled = _resolve_embeddings_flag(embeddings, notes)
|
|
98
|
+
semantic_search_payload = _import_semantic_search_payload()
|
|
99
|
+
|
|
100
|
+
query_gold = load_query_gold(fixture)
|
|
101
|
+
traverse_gold = load_traverse_gold(fixture)
|
|
102
|
+
metric_templates = load_metric_templates(fixture)
|
|
103
|
+
|
|
104
|
+
metrics: dict[str, Metric] = {}
|
|
105
|
+
per_item: list[dict[str, Any]] = []
|
|
106
|
+
|
|
107
|
+
with tempfile.TemporaryDirectory() as tmp_root:
|
|
108
|
+
tmp_dir = Path(tmp_root) / "graph"
|
|
109
|
+
gdir = build_l2_graph_dir(
|
|
110
|
+
snapshot, work_dir=tmp_dir, embeddings=embeddings_enabled
|
|
111
|
+
)
|
|
112
|
+
|
|
113
|
+
# query Recall@K --------------------------------------------------
|
|
114
|
+
query_observations = _collect_query_observations(
|
|
115
|
+
graph_dir=gdir,
|
|
116
|
+
database=snapshot.connection_name,
|
|
117
|
+
gold=query_gold,
|
|
118
|
+
)
|
|
119
|
+
per_item.extend(
|
|
120
|
+
_recall_per_item(kind="query_recall", observations=query_observations)
|
|
121
|
+
)
|
|
122
|
+
if any(g.tables_touched for g in query_gold):
|
|
123
|
+
metrics[f"query_recall_at_{_RECALL_K}"] = Metric(
|
|
124
|
+
value=query_recall_at_k(
|
|
125
|
+
[
|
|
126
|
+
(g.tables_touched, _bare_hits(hits))
|
|
127
|
+
for g, hits in query_observations
|
|
128
|
+
],
|
|
129
|
+
k=_RECALL_K,
|
|
130
|
+
),
|
|
131
|
+
direction="higher_is_better",
|
|
132
|
+
)
|
|
133
|
+
else:
|
|
134
|
+
notes.append(
|
|
135
|
+
f"query_recall_at_{_RECALL_K} skipped for dataset "
|
|
136
|
+
f"'{dataset.value}': no gold question→table mappings present "
|
|
137
|
+
"(fixture missing a *_nl2sql_bench.json with parseable SQL)."
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
# semantic_search Recall@K ---------------------------------------
|
|
141
|
+
if semantic_search_payload is None:
|
|
142
|
+
notes.append(
|
|
143
|
+
f"semantic_search_recall_at_{_RECALL_K} skipped: the "
|
|
144
|
+
"semantic_search MCP tool is not yet available in this "
|
|
145
|
+
"build. Re-run once the tool ships."
|
|
146
|
+
)
|
|
147
|
+
elif not embeddings_enabled:
|
|
148
|
+
notes.append(
|
|
149
|
+
f"semantic_search_recall_at_{_RECALL_K} skipped: "
|
|
150
|
+
"--embeddings not requested or the [embeddings] extra is "
|
|
151
|
+
"not installed."
|
|
152
|
+
)
|
|
153
|
+
else:
|
|
154
|
+
sem_observations = _collect_semantic_observations(
|
|
155
|
+
semantic_search_payload=semantic_search_payload,
|
|
156
|
+
graph_dir=gdir,
|
|
157
|
+
database=snapshot.connection_name,
|
|
158
|
+
gold=query_gold,
|
|
159
|
+
)
|
|
160
|
+
per_item.extend(
|
|
161
|
+
_recall_per_item(
|
|
162
|
+
kind="semantic_search_recall", observations=sem_observations
|
|
163
|
+
)
|
|
164
|
+
)
|
|
165
|
+
if any(g.tables_touched for g in query_gold):
|
|
166
|
+
metrics[f"semantic_search_recall_at_{_RECALL_K}"] = Metric(
|
|
167
|
+
value=semantic_search_recall_at_k(
|
|
168
|
+
[
|
|
169
|
+
(g.tables_touched, _bare_hits(hits))
|
|
170
|
+
for g, hits in sem_observations
|
|
171
|
+
],
|
|
172
|
+
k=_RECALL_K,
|
|
173
|
+
),
|
|
174
|
+
direction="higher_is_better",
|
|
175
|
+
)
|
|
176
|
+
|
|
177
|
+
# traverse correctness -------------------------------------------
|
|
178
|
+
traverse_observations = _collect_traverse_observations(
|
|
179
|
+
graph_dir=gdir,
|
|
180
|
+
database=snapshot.connection_name,
|
|
181
|
+
gold=traverse_gold,
|
|
182
|
+
)
|
|
183
|
+
per_item.extend(_traverse_per_item(traverse_observations))
|
|
184
|
+
if traverse_gold:
|
|
185
|
+
metrics["traverse_correctness"] = Metric(
|
|
186
|
+
value=traverse_correctness(
|
|
187
|
+
[(g.gold_path, paths) for g, paths in traverse_observations]
|
|
188
|
+
),
|
|
189
|
+
direction="higher_is_better",
|
|
190
|
+
)
|
|
191
|
+
else:
|
|
192
|
+
notes.append(
|
|
193
|
+
f"traverse_correctness skipped for dataset '{dataset.value}': "
|
|
194
|
+
"no gold_path entries authored in the bench JSON."
|
|
195
|
+
)
|
|
196
|
+
|
|
197
|
+
# compile_metric correctness -------------------------------------
|
|
198
|
+
cm_observations = _collect_compile_metric_observations(
|
|
199
|
+
graph_dir=gdir,
|
|
200
|
+
templates=metric_templates,
|
|
201
|
+
)
|
|
202
|
+
per_item.extend(_compile_metric_per_item(metric_templates, cm_observations))
|
|
203
|
+
if metric_templates:
|
|
204
|
+
metrics["compile_metric_correctness"] = Metric(
|
|
205
|
+
value=compile_metric_correctness(
|
|
206
|
+
[valid for valid, _err in cm_observations]
|
|
207
|
+
),
|
|
208
|
+
direction="higher_is_better",
|
|
209
|
+
)
|
|
210
|
+
else:
|
|
211
|
+
notes.append(
|
|
212
|
+
f"compile_metric_correctness skipped for dataset "
|
|
213
|
+
f"'{dataset.value}': no metric-templates YAML "
|
|
214
|
+
"(scripts/data/<dataset>_metric_templates.yaml absent)."
|
|
215
|
+
)
|
|
216
|
+
|
|
217
|
+
per_item.sort(key=lambda r: (r.get("kind", ""), r.get("id", "")))
|
|
218
|
+
|
|
219
|
+
result = BenchmarkResult(
|
|
220
|
+
level="l2",
|
|
221
|
+
dataset=dataset.value,
|
|
222
|
+
pretensor_version=_resolve_version(),
|
|
223
|
+
embeddings_enabled=embeddings_enabled,
|
|
224
|
+
ran_at=_DETERMINISTIC_RAN_AT,
|
|
225
|
+
fixture_sha=fixture_sha,
|
|
226
|
+
metrics=metrics,
|
|
227
|
+
per_item=per_item,
|
|
228
|
+
notes=notes,
|
|
229
|
+
)
|
|
230
|
+
|
|
231
|
+
if out is None:
|
|
232
|
+
sys.stdout.write(json.dumps(result.to_dict(), sort_keys=True, indent=2) + "\n")
|
|
233
|
+
else:
|
|
234
|
+
write_json(result, out)
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
# ---------------------------------------------------------------------------
|
|
238
|
+
# observation collectors
|
|
239
|
+
# ---------------------------------------------------------------------------
|
|
240
|
+
|
|
241
|
+
|
|
242
|
+
def _collect_query_observations(
|
|
243
|
+
*,
|
|
244
|
+
graph_dir: Path,
|
|
245
|
+
database: str,
|
|
246
|
+
gold: list[QueryGoldEntry],
|
|
247
|
+
) -> list[tuple[QueryGoldEntry, list[RankedHit]]]:
|
|
248
|
+
"""Run ``query_payload`` for each gold question.
|
|
249
|
+
|
|
250
|
+
The MCP tool's ``limit`` is set generously above ``_RECALL_K`` so the
|
|
251
|
+
metric's tie-aware top-K can pick up tied items past the strict cut.
|
|
252
|
+
"""
|
|
253
|
+
out: list[tuple[QueryGoldEntry, list[RankedHit]]] = []
|
|
254
|
+
fetch_limit = max(_RECALL_K * 4, 20)
|
|
255
|
+
for entry in gold:
|
|
256
|
+
payload = query_payload(
|
|
257
|
+
graph_dir, q=entry.question, db=database, limit=fetch_limit
|
|
258
|
+
)
|
|
259
|
+
ranked = _ranked_hits_from_query(payload)
|
|
260
|
+
out.append((entry, ranked))
|
|
261
|
+
return out
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
def _collect_semantic_observations(
|
|
265
|
+
*,
|
|
266
|
+
semantic_search_payload: Any,
|
|
267
|
+
graph_dir: Path,
|
|
268
|
+
database: str,
|
|
269
|
+
gold: list[QueryGoldEntry],
|
|
270
|
+
) -> list[tuple[QueryGoldEntry, list[RankedHit]]]:
|
|
271
|
+
"""Run ``semantic_search`` for each gold question.
|
|
272
|
+
|
|
273
|
+
The exact return shape depends on the tool that ships. We accept
|
|
274
|
+
any callable that returns ``{"results": [...]}`` with ``score`` and
|
|
275
|
+
``name`` fields per hit — same contract as ``query_payload`` — so
|
|
276
|
+
the forward-compatible mapping below works the moment the tool
|
|
277
|
+
lands.
|
|
278
|
+
"""
|
|
279
|
+
out: list[tuple[QueryGoldEntry, list[RankedHit]]] = []
|
|
280
|
+
fetch_limit = max(_RECALL_K * 4, 20)
|
|
281
|
+
for entry in gold:
|
|
282
|
+
payload = semantic_search_payload(
|
|
283
|
+
graph_dir,
|
|
284
|
+
query=entry.question,
|
|
285
|
+
database=database,
|
|
286
|
+
k=fetch_limit,
|
|
287
|
+
)
|
|
288
|
+
ranked = _ranked_hits_from_query(payload)
|
|
289
|
+
out.append((entry, ranked))
|
|
290
|
+
return out
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def _collect_traverse_observations(
|
|
294
|
+
*,
|
|
295
|
+
graph_dir: Path,
|
|
296
|
+
database: str,
|
|
297
|
+
gold: list[TraverseGoldEntry],
|
|
298
|
+
) -> list[tuple[TraverseGoldEntry, list[list[JoinPair]]]]:
|
|
299
|
+
"""Run ``traverse_payload`` for each gold ``(from, to)`` pair.
|
|
300
|
+
|
|
301
|
+
Returns the full list of returned paths per item — the traverse
|
|
302
|
+
tool emits all top-ranked paths when tied, and the metric counts
|
|
303
|
+
a match if ANY returned path equals the gold sequence.
|
|
304
|
+
"""
|
|
305
|
+
out: list[tuple[TraverseGoldEntry, list[list[JoinPair]]]] = []
|
|
306
|
+
for entry in gold:
|
|
307
|
+
payload = traverse_payload(
|
|
308
|
+
graph_dir,
|
|
309
|
+
from_table=entry.from_table,
|
|
310
|
+
to_table=entry.to_table,
|
|
311
|
+
database=database,
|
|
312
|
+
)
|
|
313
|
+
paths = _path_pairs_from_traverse(payload)
|
|
314
|
+
out.append((entry, paths))
|
|
315
|
+
return out
|
|
316
|
+
|
|
317
|
+
|
|
318
|
+
def _collect_compile_metric_observations(
|
|
319
|
+
*,
|
|
320
|
+
graph_dir: Path,
|
|
321
|
+
templates: list[MetricTemplateEntry],
|
|
322
|
+
) -> list[tuple[bool, str]]:
|
|
323
|
+
"""Run ``compile_metric_payload`` for each template; return ``(valid, error)``."""
|
|
324
|
+
out: list[tuple[bool, str]] = []
|
|
325
|
+
for tpl in templates:
|
|
326
|
+
payload = compile_metric_payload(
|
|
327
|
+
graph_dir,
|
|
328
|
+
semantic_yaml=tpl.semantic_yaml,
|
|
329
|
+
metric=tpl.metric,
|
|
330
|
+
database=tpl.database,
|
|
331
|
+
)
|
|
332
|
+
if "error" in payload:
|
|
333
|
+
out.append((False, str(payload["error"])))
|
|
334
|
+
else:
|
|
335
|
+
out.append((bool(payload.get("valid", False)), ""))
|
|
336
|
+
return out
|
|
337
|
+
|
|
338
|
+
|
|
339
|
+
# ---------------------------------------------------------------------------
|
|
340
|
+
# per-item record builders
|
|
341
|
+
# ---------------------------------------------------------------------------
|
|
342
|
+
|
|
343
|
+
|
|
344
|
+
def _recall_per_item(
|
|
345
|
+
*,
|
|
346
|
+
kind: str,
|
|
347
|
+
observations: list[tuple[QueryGoldEntry, list[RankedHit]]],
|
|
348
|
+
) -> list[dict[str, Any]]:
|
|
349
|
+
"""Build per-question Recall@K records for the JSON envelope.
|
|
350
|
+
|
|
351
|
+
Used for both ``query_recall`` and ``semantic_search_recall`` —
|
|
352
|
+
the only thing that varies is the ``kind`` discriminator.
|
|
353
|
+
"""
|
|
354
|
+
out: list[dict[str, Any]] = []
|
|
355
|
+
for entry, hits in observations:
|
|
356
|
+
retrieved = top_k_with_ties(hits, _RECALL_K)
|
|
357
|
+
retrieved_bare = sorted({bare_table_name(name) for _score, name in retrieved})
|
|
358
|
+
gold_bare = sorted(entry.tables_touched)
|
|
359
|
+
intersection = sorted(set(gold_bare) & set(retrieved_bare))
|
|
360
|
+
out.append(
|
|
361
|
+
{
|
|
362
|
+
"kind": kind,
|
|
363
|
+
"id": entry.id,
|
|
364
|
+
"gold_tables": gold_bare,
|
|
365
|
+
"retrieved_tables": retrieved_bare,
|
|
366
|
+
"matched": intersection,
|
|
367
|
+
"recall": (len(intersection) / len(gold_bare) if gold_bare else None),
|
|
368
|
+
}
|
|
369
|
+
)
|
|
370
|
+
return out
|
|
371
|
+
|
|
372
|
+
|
|
373
|
+
def _traverse_per_item(
|
|
374
|
+
observations: list[tuple[TraverseGoldEntry, list[list[JoinPair]]]],
|
|
375
|
+
) -> list[dict[str, Any]]:
|
|
376
|
+
out: list[dict[str, Any]] = []
|
|
377
|
+
for entry, paths in observations:
|
|
378
|
+
gold_t = tuple(tuple(p) for p in entry.gold_path)
|
|
379
|
+
match = any(tuple(tuple(p) for p in path) == gold_t for path in paths)
|
|
380
|
+
out.append(
|
|
381
|
+
{
|
|
382
|
+
"kind": "traverse",
|
|
383
|
+
"id": entry.id,
|
|
384
|
+
"from_table": entry.from_table,
|
|
385
|
+
"to_table": entry.to_table,
|
|
386
|
+
"gold_path": [list(hop) for hop in entry.gold_path],
|
|
387
|
+
"returned_paths": [[list(hop) for hop in path] for path in paths],
|
|
388
|
+
"correct": match,
|
|
389
|
+
}
|
|
390
|
+
)
|
|
391
|
+
return out
|
|
392
|
+
|
|
393
|
+
|
|
394
|
+
def _compile_metric_per_item(
|
|
395
|
+
templates: list[MetricTemplateEntry],
|
|
396
|
+
observations: list[tuple[bool, str]],
|
|
397
|
+
) -> list[dict[str, Any]]:
|
|
398
|
+
out: list[dict[str, Any]] = []
|
|
399
|
+
for tpl, (valid, error) in zip(templates, observations, strict=True):
|
|
400
|
+
out.append(
|
|
401
|
+
{
|
|
402
|
+
"kind": "compile_metric",
|
|
403
|
+
"id": tpl.metric,
|
|
404
|
+
"valid": valid,
|
|
405
|
+
"error": error or None,
|
|
406
|
+
}
|
|
407
|
+
)
|
|
408
|
+
return out
|
|
409
|
+
|
|
410
|
+
|
|
411
|
+
# ---------------------------------------------------------------------------
|
|
412
|
+
# extractors
|
|
413
|
+
# ---------------------------------------------------------------------------
|
|
414
|
+
|
|
415
|
+
|
|
416
|
+
def _ranked_hits_from_query(payload: dict[str, Any]) -> list[RankedHit]:
|
|
417
|
+
"""Pull ``[(score, name), ...]`` out of a ``query_payload`` response."""
|
|
418
|
+
hits: list[RankedHit] = []
|
|
419
|
+
if not isinstance(payload, dict):
|
|
420
|
+
return hits
|
|
421
|
+
results = payload.get("results")
|
|
422
|
+
if not isinstance(results, list):
|
|
423
|
+
return hits
|
|
424
|
+
for hit in results:
|
|
425
|
+
if not isinstance(hit, dict):
|
|
426
|
+
continue
|
|
427
|
+
name = hit.get("name")
|
|
428
|
+
score = hit.get("score")
|
|
429
|
+
if isinstance(name, str) and isinstance(score, (int, float)):
|
|
430
|
+
hits.append((float(score), name))
|
|
431
|
+
return hits
|
|
432
|
+
|
|
433
|
+
|
|
434
|
+
def _path_pairs_from_traverse(
|
|
435
|
+
payload: dict[str, Any],
|
|
436
|
+
) -> list[list[JoinPair]]:
|
|
437
|
+
"""Pull ``[[(from_table, to_table), ...], ...]`` from a traverse response."""
|
|
438
|
+
out: list[list[JoinPair]] = []
|
|
439
|
+
if not isinstance(payload, dict) or "error" in payload:
|
|
440
|
+
return out
|
|
441
|
+
paths = payload.get("paths")
|
|
442
|
+
if not isinstance(paths, list):
|
|
443
|
+
return out
|
|
444
|
+
for path in paths:
|
|
445
|
+
if not isinstance(path, dict):
|
|
446
|
+
continue
|
|
447
|
+
steps = path.get("steps")
|
|
448
|
+
if not isinstance(steps, list):
|
|
449
|
+
continue
|
|
450
|
+
hops: list[JoinPair] = []
|
|
451
|
+
for step in steps:
|
|
452
|
+
if not isinstance(step, dict):
|
|
453
|
+
continue
|
|
454
|
+
frm = step.get("from_table")
|
|
455
|
+
to = step.get("to_table")
|
|
456
|
+
if isinstance(frm, str) and isinstance(to, str):
|
|
457
|
+
hops.append((frm, to))
|
|
458
|
+
if hops:
|
|
459
|
+
out.append(hops)
|
|
460
|
+
return out
|
|
461
|
+
|
|
462
|
+
|
|
463
|
+
def _bare_hits(hits: list[RankedHit]) -> list[RankedHit]:
|
|
464
|
+
"""Strip the schema prefix from each hit's table name.
|
|
465
|
+
|
|
466
|
+
The metric's gold side carries bare table names (the gold loader
|
|
467
|
+
strips schema prefixes when populating ``tables_touched``). The
|
|
468
|
+
MCP ``query`` payload returns ``schema.table``; we collapse before
|
|
469
|
+
the intersection so the recall calculation lines up.
|
|
470
|
+
"""
|
|
471
|
+
return [(score, bare_table_name(name)) for score, name in hits]
|
|
472
|
+
|
|
473
|
+
|
|
474
|
+
# ---------------------------------------------------------------------------
|
|
475
|
+
# helpers
|
|
476
|
+
# ---------------------------------------------------------------------------
|
|
477
|
+
|
|
478
|
+
|
|
479
|
+
def _resolve_embeddings_flag(requested: bool, notes: list[str]) -> bool:
|
|
480
|
+
"""Mirror L1's resolver: probe ``onnxruntime`` for the [embeddings] extra."""
|
|
481
|
+
if not requested:
|
|
482
|
+
return False
|
|
483
|
+
try:
|
|
484
|
+
importlib.import_module(_EMBEDDINGS_SENTINEL)
|
|
485
|
+
except ImportError:
|
|
486
|
+
notes.append(
|
|
487
|
+
"--embeddings was requested but the 'embeddings' extra is not "
|
|
488
|
+
"installed; the semantic_search metric will be omitted. Install "
|
|
489
|
+
"with `uv sync --extra embeddings` to enable it."
|
|
490
|
+
)
|
|
491
|
+
return False
|
|
492
|
+
return True
|
|
493
|
+
|
|
494
|
+
|
|
495
|
+
def _import_semantic_search_payload() -> Any:
|
|
496
|
+
"""Forward-compatible import of the semantic_search MCP tool.
|
|
497
|
+
|
|
498
|
+
Returns the ``semantic_search_payload`` callable when the module
|
|
499
|
+
exists, or ``None`` when the module itself has not landed yet.
|
|
500
|
+
This lets L2 keep its public contract ("emit
|
|
501
|
+
``semantic_search_recall_at_5`` when embeddings are installed")
|
|
502
|
+
satisfiable the moment the tool ships, with no L2 code change.
|
|
503
|
+
|
|
504
|
+
Raises:
|
|
505
|
+
AttributeError: when the module loads but does not expose
|
|
506
|
+
``semantic_search_payload``. We deliberately do NOT swallow
|
|
507
|
+
this — a future rename of the entry point should fail
|
|
508
|
+
loudly here, not silently skip the metric.
|
|
509
|
+
"""
|
|
510
|
+
try:
|
|
511
|
+
module = importlib.import_module("pretensor.mcp.tools.semantic_search")
|
|
512
|
+
except ImportError:
|
|
513
|
+
return None
|
|
514
|
+
try:
|
|
515
|
+
return module.semantic_search_payload
|
|
516
|
+
except AttributeError as exc:
|
|
517
|
+
raise AttributeError(
|
|
518
|
+
"pretensor.mcp.tools.semantic_search loaded but does not "
|
|
519
|
+
"expose `semantic_search_payload`; the L2 runner expects "
|
|
520
|
+
"this entry point. If the tool was renamed, update L2 to "
|
|
521
|
+
"match — silently skipping the metric here would mask the "
|
|
522
|
+
"regression."
|
|
523
|
+
) from exc
|
|
524
|
+
|
|
525
|
+
|
|
526
|
+
def _resolve_version() -> str:
|
|
527
|
+
try:
|
|
528
|
+
return importlib.metadata.version("pretensor")
|
|
529
|
+
except importlib.metadata.PackageNotFoundError:
|
|
530
|
+
return "0.0.0+unknown"
|
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
"""L3 benchmark runners — agent task success.
|
|
2
|
+
|
|
3
|
+
L3 measures end-to-end NL-to-SQL against a real database with two control
|
|
4
|
+
conditions:
|
|
5
|
+
|
|
6
|
+
* ``baseline`` — the agent receives the raw schema DDL only.
|
|
7
|
+
* ``pretensor`` — the agent receives a live ``pretensor serve`` MCP
|
|
8
|
+
session against an indexed graph of the same database.
|
|
9
|
+
|
|
10
|
+
This package is LLM-gated and lives outside the OSS deterministic core
|
|
11
|
+
(``docs/contracts/architecture.md`` Invariant #6). All LLM calls go through
|
|
12
|
+
``httpx`` against provider HTTP endpoints — no SDK dependency.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
from __future__ import annotations
|
|
16
|
+
|
|
17
|
+
from pretensor.benchmark.l3.agent import (
|
|
18
|
+
DEFAULT_MAX_ITERATIONS,
|
|
19
|
+
AgentLlmClient,
|
|
20
|
+
AgentLoopError,
|
|
21
|
+
AgentLoopResult,
|
|
22
|
+
AgentMessage,
|
|
23
|
+
AgentStep,
|
|
24
|
+
AgentTool,
|
|
25
|
+
AgentToolCall,
|
|
26
|
+
AgentToolResult,
|
|
27
|
+
ToolCallTraceEntry,
|
|
28
|
+
ToolInvocationOutcome,
|
|
29
|
+
ToolInvoker,
|
|
30
|
+
run_agent_loop,
|
|
31
|
+
)
|
|
32
|
+
from pretensor.benchmark.l3.llm_client import (
|
|
33
|
+
AnthropicHttpClient,
|
|
34
|
+
LlmCallError,
|
|
35
|
+
LlmClient,
|
|
36
|
+
LlmResponse,
|
|
37
|
+
OpenAIHttpClient,
|
|
38
|
+
)
|
|
39
|
+
from pretensor.benchmark.l3.mcp_client import (
|
|
40
|
+
McpClient,
|
|
41
|
+
McpClientError,
|
|
42
|
+
McpToolResult,
|
|
43
|
+
StdioMcpClient,
|
|
44
|
+
)
|
|
45
|
+
from pretensor.benchmark.l3.pretensor_runner import run_l3_pretensor
|
|
46
|
+
from pretensor.benchmark.l3.runner import run_l3_baseline
|
|
47
|
+
|
|
48
|
+
__all__ = [
|
|
49
|
+
"DEFAULT_MAX_ITERATIONS",
|
|
50
|
+
"AgentLlmClient",
|
|
51
|
+
"AgentLoopError",
|
|
52
|
+
"AgentLoopResult",
|
|
53
|
+
"AgentMessage",
|
|
54
|
+
"AgentStep",
|
|
55
|
+
"AgentTool",
|
|
56
|
+
"AgentToolCall",
|
|
57
|
+
"AgentToolResult",
|
|
58
|
+
"AnthropicHttpClient",
|
|
59
|
+
"LlmCallError",
|
|
60
|
+
"LlmClient",
|
|
61
|
+
"LlmResponse",
|
|
62
|
+
"McpClient",
|
|
63
|
+
"McpClientError",
|
|
64
|
+
"McpToolResult",
|
|
65
|
+
"OpenAIHttpClient",
|
|
66
|
+
"StdioMcpClient",
|
|
67
|
+
"ToolCallTraceEntry",
|
|
68
|
+
"ToolInvocationOutcome",
|
|
69
|
+
"ToolInvoker",
|
|
70
|
+
"run_agent_loop",
|
|
71
|
+
"run_l3_baseline",
|
|
72
|
+
"run_l3_pretensor",
|
|
73
|
+
]
|