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.
Files changed (198) hide show
  1. pretensor/__init__.py +50 -0
  2. pretensor/benchmark/__init__.py +54 -0
  3. pretensor/benchmark/cli.py +294 -0
  4. pretensor/benchmark/fixtures.py +84 -0
  5. pretensor/benchmark/l1/__init__.py +23 -0
  6. pretensor/benchmark/l1/metrics.py +141 -0
  7. pretensor/benchmark/l1/pipeline.py +188 -0
  8. pretensor/benchmark/l1/runner.py +245 -0
  9. pretensor/benchmark/l2/__init__.py +27 -0
  10. pretensor/benchmark/l2/gold.py +236 -0
  11. pretensor/benchmark/l2/metrics.py +146 -0
  12. pretensor/benchmark/l2/pipeline.py +124 -0
  13. pretensor/benchmark/l2/runner.py +530 -0
  14. pretensor/benchmark/l3/__init__.py +73 -0
  15. pretensor/benchmark/l3/agent.py +316 -0
  16. pretensor/benchmark/l3/db.py +188 -0
  17. pretensor/benchmark/l3/gold.py +85 -0
  18. pretensor/benchmark/l3/llm_client.py +395 -0
  19. pretensor/benchmark/l3/mcp_client.py +357 -0
  20. pretensor/benchmark/l3/pretensor_runner.py +456 -0
  21. pretensor/benchmark/l3/prompt.py +132 -0
  22. pretensor/benchmark/l3/runner.py +358 -0
  23. pretensor/benchmark/l3/sql_equivalence.py +176 -0
  24. pretensor/benchmark/release_gate.py +448 -0
  25. pretensor/benchmark/results.py +298 -0
  26. pretensor/benchmark/runner.py +109 -0
  27. pretensor/cli/__init__.py +1 -0
  28. pretensor/cli/commands/_source_runner.py +147 -0
  29. pretensor/cli/commands/analyze.py +201 -0
  30. pretensor/cli/commands/connections/__init__.py +7 -0
  31. pretensor/cli/commands/connections/add_remove.py +126 -0
  32. pretensor/cli/commands/connections/register.py +12 -0
  33. pretensor/cli/commands/export.py +131 -0
  34. pretensor/cli/commands/index.py +559 -0
  35. pretensor/cli/commands/list.py +76 -0
  36. pretensor/cli/commands/quickstart.py +207 -0
  37. pretensor/cli/commands/reindex.py +646 -0
  38. pretensor/cli/commands/semantic.py +190 -0
  39. pretensor/cli/commands/serve.py +144 -0
  40. pretensor/cli/commands/sync_grants.py +149 -0
  41. pretensor/cli/commands/validate.py +176 -0
  42. pretensor/cli/config_file.py +442 -0
  43. pretensor/cli/constants.py +10 -0
  44. pretensor/cli/dbt_enrichment.py +96 -0
  45. pretensor/cli/main.py +109 -0
  46. pretensor/cli/paths.py +43 -0
  47. pretensor/cli/plugin.py +52 -0
  48. pretensor/config.py +226 -0
  49. pretensor/connectors/__init__.py +29 -0
  50. pretensor/connectors/base.py +165 -0
  51. pretensor/connectors/bigquery.py +468 -0
  52. pretensor/connectors/inspect.py +321 -0
  53. pretensor/connectors/lineage_sqlglot.py +97 -0
  54. pretensor/connectors/models.py +130 -0
  55. pretensor/connectors/mysql.py +402 -0
  56. pretensor/connectors/pg_array_parse.py +53 -0
  57. pretensor/connectors/postgres.py +938 -0
  58. pretensor/connectors/registry.py +93 -0
  59. pretensor/connectors/snapshot.py +244 -0
  60. pretensor/connectors/snowflake.py +908 -0
  61. pretensor/core/__init__.py +1 -0
  62. pretensor/core/builder.py +307 -0
  63. pretensor/core/dsn_crypto.py +51 -0
  64. pretensor/core/graph_schema_manager.py +246 -0
  65. pretensor/core/graph_store.py +1226 -0
  66. pretensor/core/ids.py +101 -0
  67. pretensor/core/portable_export.py +276 -0
  68. pretensor/core/query_runner.py +67 -0
  69. pretensor/core/registry.py +209 -0
  70. pretensor/core/schema.py +473 -0
  71. pretensor/core/secure_io.py +93 -0
  72. pretensor/core/store.py +469 -0
  73. pretensor/enrichment/__init__.py +1 -0
  74. pretensor/enrichment/analyze/__init__.py +0 -0
  75. pretensor/enrichment/analyze/classify.py +49 -0
  76. pretensor/enrichment/analyze/extract_python.py +196 -0
  77. pretensor/enrichment/analyze/parse.py +141 -0
  78. pretensor/enrichment/analyze/pipeline.py +195 -0
  79. pretensor/enrichment/analyze/summary.py +38 -0
  80. pretensor/enrichment/analyze/walker.py +98 -0
  81. pretensor/enrichment/analyze/writers.py +214 -0
  82. pretensor/enrichment/dbt/__init__.py +30 -0
  83. pretensor/enrichment/dbt/lineage.py +100 -0
  84. pretensor/enrichment/dbt/manifest.py +300 -0
  85. pretensor/enrichment/dbt/metadata.py +263 -0
  86. pretensor/enrichment/dbt/pipeline.py +77 -0
  87. pretensor/enrichment/dbt/resolution.py +101 -0
  88. pretensor/enrichment/dbt/signals.py +305 -0
  89. pretensor/entities/__init__.py +27 -0
  90. pretensor/entities/builder.py +63 -0
  91. pretensor/entities/classifier.py +383 -0
  92. pretensor/entities/llm_extract.py +66 -0
  93. pretensor/errors.py +35 -0
  94. pretensor/graph_models/__init__.py +17 -0
  95. pretensor/graph_models/base.py +11 -0
  96. pretensor/graph_models/consumer.py +71 -0
  97. pretensor/graph_models/edge.py +35 -0
  98. pretensor/graph_models/entity.py +21 -0
  99. pretensor/graph_models/node.py +79 -0
  100. pretensor/graph_models/relationship.py +33 -0
  101. pretensor/integrations/__init__.py +42 -0
  102. pretensor/integrations/_base.py +138 -0
  103. pretensor/integrations/google_adk.py +49 -0
  104. pretensor/integrations/langchain.py +55 -0
  105. pretensor/integrations/llamaindex.py +53 -0
  106. pretensor/intelligence/__init__.py +33 -0
  107. pretensor/intelligence/cluster_labeler.py +425 -0
  108. pretensor/intelligence/clustering.py +168 -0
  109. pretensor/intelligence/combining.py +32 -0
  110. pretensor/intelligence/discovery.py +114 -0
  111. pretensor/intelligence/embeddings.py +317 -0
  112. pretensor/intelligence/graph_export.py +200 -0
  113. pretensor/intelligence/heuristic.py +544 -0
  114. pretensor/intelligence/join_paths/__init__.py +130 -0
  115. pretensor/intelligence/join_paths/on_demand.py +516 -0
  116. pretensor/intelligence/join_paths/storage.py +70 -0
  117. pretensor/intelligence/llm_infer.py +78 -0
  118. pretensor/intelligence/llm_runtime.py +62 -0
  119. pretensor/intelligence/metric_templates.py +193 -0
  120. pretensor/intelligence/pipeline.py +364 -0
  121. pretensor/intelligence/role_exemplars.py +263 -0
  122. pretensor/intelligence/schema_classification.py +360 -0
  123. pretensor/intelligence/scoring.py +76 -0
  124. pretensor/intelligence/semantic.py +240 -0
  125. pretensor/intelligence/shadow_alias.py +101 -0
  126. pretensor/intelligence/statistical.py +50 -0
  127. pretensor/intelligence/steps.py +191 -0
  128. pretensor/intelligence/steps_embedding.py +168 -0
  129. pretensor/introspection/__init__.py +6 -0
  130. pretensor/introspection/inspector.py +5 -0
  131. pretensor/introspection/models/__init__.py +0 -0
  132. pretensor/introspection/models/base.py +5 -0
  133. pretensor/introspection/models/config.py +237 -0
  134. pretensor/introspection/models/dsn.py +550 -0
  135. pretensor/introspection/models/plan.py +116 -0
  136. pretensor/introspection/models/schema.py +10 -0
  137. pretensor/introspection/models/semantic.py +121 -0
  138. pretensor/introspection/models/validation.py +116 -0
  139. pretensor/introspection/snapshot.py +46 -0
  140. pretensor/mcp/__init__.py +16 -0
  141. pretensor/mcp/config_json.py +24 -0
  142. pretensor/mcp/payload_types.py +274 -0
  143. pretensor/mcp/resources/__init__.py +17 -0
  144. pretensor/mcp/resources/markdown.py +314 -0
  145. pretensor/mcp/server.py +285 -0
  146. pretensor/mcp/service.py +49 -0
  147. pretensor/mcp/service_context.py +142 -0
  148. pretensor/mcp/service_registry.py +294 -0
  149. pretensor/mcp/store_cache.py +43 -0
  150. pretensor/mcp/tool_registry.py +136 -0
  151. pretensor/mcp/tools/__init__.py +1 -0
  152. pretensor/mcp/tools/_rank.py +244 -0
  153. pretensor/mcp/tools/_timed.py +26 -0
  154. pretensor/mcp/tools/compile_metric.py +144 -0
  155. pretensor/mcp/tools/consumers.py +161 -0
  156. pretensor/mcp/tools/context.py +1121 -0
  157. pretensor/mcp/tools/cypher.py +509 -0
  158. pretensor/mcp/tools/detect_changes.py +254 -0
  159. pretensor/mcp/tools/impact.py +271 -0
  160. pretensor/mcp/tools/list.py +131 -0
  161. pretensor/mcp/tools/schema.py +170 -0
  162. pretensor/mcp/tools/search.py +316 -0
  163. pretensor/mcp/tools/semantic_search.py +282 -0
  164. pretensor/mcp/tools/traverse.py +1027 -0
  165. pretensor/mcp/tools/validate_sql.py +150 -0
  166. pretensor/observability.py +203 -0
  167. pretensor/py.typed +0 -0
  168. pretensor/quickstart/README.md +29 -0
  169. pretensor/quickstart/__init__.py +6 -0
  170. pretensor/quickstart/docker-compose.yml +18 -0
  171. pretensor/quickstart/pagila_data.sql +63 -0
  172. pretensor/quickstart/pagila_ddl.sql +92 -0
  173. pretensor/search/__init__.py +6 -0
  174. pretensor/search/base.py +80 -0
  175. pretensor/search/index.py +435 -0
  176. pretensor/semantic/__init__.py +24 -0
  177. pretensor/semantic/base.py +123 -0
  178. pretensor/semantic/compiler.py +487 -0
  179. pretensor/semantic/yaml_layer.py +180 -0
  180. pretensor/skills/__init__.py +5 -0
  181. pretensor/skills/generator.py +235 -0
  182. pretensor/staleness/__init__.py +15 -0
  183. pretensor/staleness/graph_patcher.py +355 -0
  184. pretensor/staleness/impact_analyzer.py +162 -0
  185. pretensor/staleness/snapshot_store.py +38 -0
  186. pretensor/validation/__init__.py +9 -0
  187. pretensor/validation/query_validator.py +436 -0
  188. pretensor/visibility/__init__.py +23 -0
  189. pretensor/visibility/config.py +126 -0
  190. pretensor/visibility/filter.py +143 -0
  191. pretensor/visibility/kuzu_helpers.py +32 -0
  192. pretensor/visibility/runtime.py +36 -0
  193. pretensor/visibility/sync_grants.py +188 -0
  194. pretensor-0.1.0.dist-info/METADATA +251 -0
  195. pretensor-0.1.0.dist-info/RECORD +198 -0
  196. pretensor-0.1.0.dist-info/WHEEL +4 -0
  197. pretensor-0.1.0.dist-info/entry_points.txt +2 -0
  198. pretensor-0.1.0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,188 @@
1
+ """Pipeline orchestration for L1 metric collection.
2
+
3
+ Glues the existing :class:`GraphBuilder` + :class:`KuzuStore` +
4
+ :class:`ClusteringEngine` + :func:`classify_database_tables` plumbing
5
+ into a single helper that returns the artifacts the four L1 metrics
6
+ need. Lives separately from :mod:`pretensor.benchmark.l1.metrics` so
7
+ the pure metric helpers stay testable without spinning up Kuzu.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from dataclasses import dataclass, field
13
+ from pathlib import Path
14
+
15
+ from pretensor.benchmark.l1.metrics import JoinKey
16
+ from pretensor.connectors.models import SchemaSnapshot
17
+ from pretensor.core.builder import GraphBuilder
18
+ from pretensor.core.store import KuzuStore
19
+ from pretensor.intelligence.clustering import ClusteringEngine
20
+ from pretensor.intelligence.discovery import RelationshipDiscovery
21
+ from pretensor.intelligence.graph_export import GraphExporter
22
+ from pretensor.intelligence.heuristic import HeuristicScorer
23
+ from pretensor.intelligence.schema_classification import classify_database_tables
24
+ from pretensor.intelligence.scoring import ScorerRegistry
25
+
26
+ __all__ = [
27
+ "L1_EMBEDDING_JOIN_THRESHOLD",
28
+ "L1Artifacts",
29
+ "build_l1_artifacts",
30
+ "discover_inferred_joins_blind",
31
+ ]
32
+
33
+ # Cosine threshold used by the L1 embeddings lane for the embedding
34
+ # relationship scorer. ``EmbeddingsConfig.join_threshold`` has no
35
+ # production default (None = scorer off); the benchmark pins the
36
+ # documented "typical" value so the lane measures one fixed, reproducible
37
+ # operating point. Bump deliberately and regenerate the
38
+ # ``<dataset>-l1-embeddings.json`` baselines when retuning.
39
+ L1_EMBEDDING_JOIN_THRESHOLD = 0.85
40
+
41
+
42
+ @dataclass(frozen=True, slots=True)
43
+ class L1Artifacts:
44
+ """Per-run outputs collected from the intelligence pipeline."""
45
+
46
+ clusters: list[frozenset[str]] = field(default_factory=list)
47
+ """One frozenset of table node IDs per cluster."""
48
+
49
+ roles: dict[str, str] = field(default_factory=dict)
50
+ """Map of bare table name → role literal (matches ``TableRole``)."""
51
+
52
+
53
+ def build_l1_artifacts(
54
+ snapshot: SchemaSnapshot,
55
+ *,
56
+ work_dir: Path,
57
+ ) -> L1Artifacts:
58
+ """Run the production pipeline once and collect L1-relevant outputs.
59
+
60
+ The pipeline runs over the snapshot as authored, with declared FKs
61
+ in place — clustering and role classification both take advantage
62
+ of that signal in production, and the L1 metric must reflect that.
63
+
64
+ The blind-discovery leg (FKs masked, used for inferred-join P/R)
65
+ runs separately via :func:`discover_inferred_joins_blind` so the
66
+ two legs don't share state.
67
+ """
68
+ work_dir.mkdir(parents=True, exist_ok=True)
69
+ store = KuzuStore(work_dir / "graph.kuzu")
70
+ try:
71
+ GraphBuilder().build(snapshot, store)
72
+
73
+ graph = GraphExporter(store).to_igraph(snapshot.database)
74
+ clusters_raw = ClusteringEngine().cluster(graph)
75
+ clusters = sorted(
76
+ (frozenset(c.table_ids) for c in clusters_raw),
77
+ key=lambda fs: tuple(sorted(fs)),
78
+ )
79
+
80
+ classifications = classify_database_tables(store, snapshot.database)
81
+ roles = {
82
+ _bare_table_name(node_id, snapshot.connection_name): cls.role
83
+ for node_id, cls in classifications.items()
84
+ }
85
+ finally:
86
+ store.close()
87
+
88
+ return L1Artifacts(clusters=list(clusters), roles=roles)
89
+
90
+
91
+ def discover_inferred_joins_blind(
92
+ snapshot: SchemaSnapshot,
93
+ *,
94
+ work_dir: Path,
95
+ embeddings: bool = False,
96
+ ) -> list[JoinKey]:
97
+ """Run heuristic discovery against an FK-masked snapshot.
98
+
99
+ Production discovery filters out candidates that match declared FKs,
100
+ so a metric that compares the inferred set against declared FKs only
101
+ makes sense when those FKs are hidden from the heuristic. This
102
+ helper masks the snapshot, runs the heuristic, and returns the
103
+ candidate joins as canonical (table-name, column) pairs ready for
104
+ :func:`pretensor.benchmark.l1.metrics.inferred_join_pr`.
105
+
106
+ With ``embeddings=True`` (the L1 embeddings lane), table vectors are
107
+ computed on the masked store first and an
108
+ :class:`EmbeddingRelationshipScorer` is registered after the heuristic
109
+ at :data:`L1_EMBEDDING_JOIN_THRESHOLD` — the same staging production
110
+ uses — so ``inferred_join_precision`` / ``recall`` measure the
111
+ heuristic + embedding scorer combination.
112
+ """
113
+ work_dir.mkdir(parents=True, exist_ok=True)
114
+ masked = _mask_foreign_keys(snapshot)
115
+ store = KuzuStore(work_dir / "graph.kuzu")
116
+ try:
117
+ # Build the FK-masked graph so RelationshipDiscovery has node
118
+ # context (column metadata, neighbours) to score candidates.
119
+ GraphBuilder().build(masked, store, run_relationship_discovery=False)
120
+ scorers = None
121
+ if embeddings:
122
+ from pretensor.intelligence.semantic import (
123
+ extend_with_embedding_scorer,
124
+ )
125
+ from pretensor.intelligence.steps_embedding import (
126
+ compute_table_embeddings,
127
+ )
128
+
129
+ compute_table_embeddings(store, masked.database)
130
+ if not store.has_any_table_embeddings():
131
+ # Mirror the L2 guard: the production embed path degrades
132
+ # silently on download failure, but a benchmark lane that
133
+ # measures the heuristic-only path under an embeddings
134
+ # baseline reports a bogus regression. Fail loudly.
135
+ msg = (
136
+ "L1 was invoked with --embeddings but no table vectors "
137
+ "were computed; the embedding model is likely "
138
+ "unavailable (download failure / rate limit)."
139
+ )
140
+ raise RuntimeError(msg)
141
+ scorers = extend_with_embedding_scorer(
142
+ ScorerRegistry([HeuristicScorer()]),
143
+ store=store,
144
+ database_key=masked.database,
145
+ threshold=L1_EMBEDDING_JOIN_THRESHOLD,
146
+ )
147
+ candidates = RelationshipDiscovery(store, scorers=scorers).discover(masked)
148
+ finally:
149
+ store.close()
150
+
151
+ keys: list[JoinKey] = []
152
+ for cand in candidates:
153
+ src = _bare_table_name(cand.source_node_id, snapshot.connection_name)
154
+ dst = _bare_table_name(cand.target_node_id, snapshot.connection_name)
155
+ keys.append((src, cand.source_column, dst, cand.target_column))
156
+ keys.sort()
157
+ return keys
158
+
159
+
160
+ def collect_declared_fks(snapshot: SchemaSnapshot) -> list[JoinKey]:
161
+ """Return the canonical declared-FK ground-truth set for the snapshot."""
162
+ keys: list[JoinKey] = []
163
+ for table in snapshot.tables:
164
+ for fk in table.foreign_keys:
165
+ src = f"{fk.source_schema}.{fk.source_table}"
166
+ dst = f"{fk.target_schema}.{fk.target_table}"
167
+ keys.append((src, fk.source_column, dst, fk.target_column))
168
+ keys.sort()
169
+ return keys
170
+
171
+
172
+ def _mask_foreign_keys(snapshot: SchemaSnapshot) -> SchemaSnapshot:
173
+ """Return a copy of ``snapshot`` with every table's ``foreign_keys`` cleared.
174
+
175
+ Used so :class:`RelationshipDiscovery` has to rediscover the joins
176
+ from column-level signals rather than reading them off the FK list.
177
+ """
178
+ masked_tables = []
179
+ for table in snapshot.tables:
180
+ masked_tables.append(table.model_copy(update={"foreign_keys": []}))
181
+ return snapshot.model_copy(update={"tables": masked_tables})
182
+
183
+
184
+ def _bare_table_name(node_id: str, connection_name: str) -> str:
185
+ """Turn a Kuzu ``node_id`` (``conn::schema::table``) into ``schema.table``."""
186
+ prefix = f"{connection_name}::"
187
+ rest = node_id[len(prefix) :] if node_id.startswith(prefix) else node_id
188
+ return rest.replace("::", ".", 1)
@@ -0,0 +1,245 @@
1
+ """``run_l1`` — collect L1 metrics into a deterministic JSON document."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import importlib
7
+ import importlib.metadata
8
+ import json
9
+ import sys
10
+ import tempfile
11
+ from pathlib import Path
12
+ from typing import TYPE_CHECKING
13
+
14
+ import yaml
15
+
16
+ from pretensor.benchmark.fixtures import load_dataset
17
+ from pretensor.benchmark.l1.metrics import (
18
+ canonicalise_join_key,
19
+ cluster_stability_jaccard,
20
+ inferred_join_pr,
21
+ role_f1,
22
+ )
23
+ from pretensor.benchmark.l1.pipeline import (
24
+ L1_EMBEDDING_JOIN_THRESHOLD,
25
+ build_l1_artifacts,
26
+ collect_declared_fks,
27
+ discover_inferred_joins_blind,
28
+ )
29
+ from pretensor.benchmark.results import BenchmarkResult, Metric, write_json
30
+ from pretensor.connectors.models import SchemaSnapshot
31
+
32
+ if TYPE_CHECKING:
33
+ from pretensor.benchmark.runner import Dataset
34
+
35
+ __all__ = ["run_l1"]
36
+
37
+
38
+ _DETERMINISTIC_RAN_AT = "1970-01-01T00:00:00Z"
39
+ """Pinned timestamp so two consecutive runs produce byte-identical JSON.
40
+
41
+ Matches the precedent in ``tests/benchmark/test_results.py``. Real
42
+ wall-clock would break the determinism gate baked into the spec.
43
+ """
44
+
45
+ _EMBEDDINGS_SENTINEL = "onnxruntime"
46
+ """Importable module that ships with the ``[embeddings]`` extra.
47
+
48
+ The runner uses this as a probe: if ``import onnxruntime`` fails the
49
+ extra is not installed, and we report the install hint rather than
50
+ silently emitting heuristic-only numbers as if embeddings were on.
51
+ """
52
+
53
+ _FIXTURES_ROOT = Path(__file__).resolve().parents[3].parent / "tests" / "fixtures"
54
+ """``<repo>/tests/fixtures`` resolved relative to the source tree.
55
+
56
+ Used to locate gold-role companion files (``<dataset>_roles.yaml``).
57
+ ``parents[3]`` walks up ``runner.py → l1 → benchmark → pretensor`` and
58
+ ``.parent`` strips ``src/``; the result is the repo root.
59
+ """
60
+
61
+
62
+ def run_l1(
63
+ dataset: Dataset,
64
+ out: Path | None,
65
+ graph_dir: Path, # noqa: ARG001 — L1 builds a fresh in-process graph
66
+ *,
67
+ embeddings: bool,
68
+ ) -> None:
69
+ """Run the four OSS L1 metrics for ``dataset`` and emit a ``BenchmarkResult``.
70
+
71
+ ``graph_dir`` is part of the CLI contract for L2 / L3 but unused here —
72
+ L1 indexes the fixture into a fresh per-call temporary store so the
73
+ metric is reproducible from a clean checkout. It's accepted (and
74
+ ignored) to keep ``run_l1`` interchangeable with the other runners.
75
+ """
76
+ fixture = load_dataset(dataset)
77
+ snapshot_text = fixture.schema_yaml_path.read_text(encoding="utf-8")
78
+ snapshot = SchemaSnapshot.from_yaml(snapshot_text)
79
+ fixture_sha = (
80
+ "sha256:" + hashlib.sha256(fixture.schema_yaml_path.read_bytes()).hexdigest()
81
+ )
82
+
83
+ notes: list[str] = []
84
+ embeddings_enabled = _resolve_embeddings_flag(embeddings, notes)
85
+
86
+ if embeddings_enabled:
87
+ notes.append(
88
+ "inferred_join metrics include the embedding relationship "
89
+ f"scorer at cosine threshold {L1_EMBEDDING_JOIN_THRESHOLD} "
90
+ "(heuristic + embedding lane)."
91
+ )
92
+
93
+ with tempfile.TemporaryDirectory() as tmp_root:
94
+ tmp = Path(tmp_root)
95
+ artifacts_a = build_l1_artifacts(snapshot, work_dir=tmp / "run-a")
96
+ artifacts_b = build_l1_artifacts(snapshot, work_dir=tmp / "run-b")
97
+ inferred = discover_inferred_joins_blind(
98
+ snapshot, work_dir=tmp / "blind", embeddings=embeddings_enabled
99
+ )
100
+
101
+ declared_fks = collect_declared_fks(snapshot)
102
+ p, r = inferred_join_pr(inferred, declared_fks)
103
+ jacc = cluster_stability_jaccard(artifacts_a.clusters, artifacts_b.clusters)
104
+
105
+ gold_roles = _load_gold_roles(dataset)
106
+ if gold_roles is None:
107
+ rf: float | None = None
108
+ notes.append(
109
+ f"role_f1 unavailable for dataset '{dataset.value}': "
110
+ f"no gold-role annotation file (only adversarial has labels today)."
111
+ )
112
+ else:
113
+ # Normalise predicted keys (``schema.table``) to bare table names so
114
+ # they line up with the gold-role file's bare-name keys.
115
+ predicted_bare: dict[str, str] = {}
116
+ for key, role in artifacts_a.roles.items():
117
+ bare = key.split(".", 1)[1] if "." in key else key
118
+ predicted_bare[bare] = role
119
+ rf = role_f1(predicted_bare, gold_roles)
120
+
121
+ metrics = {
122
+ "inferred_join_precision": Metric(value=p, direction="higher_is_better"),
123
+ "inferred_join_recall": Metric(value=r, direction="higher_is_better"),
124
+ "cluster_stability_jaccard": Metric(value=jacc, direction="higher_is_better"),
125
+ "role_f1": Metric(value=rf, direction="higher_is_better"),
126
+ }
127
+
128
+ result = BenchmarkResult(
129
+ level="l1",
130
+ dataset=dataset.value,
131
+ pretensor_version=_resolve_version(),
132
+ embeddings_enabled=embeddings_enabled,
133
+ ran_at=_DETERMINISTIC_RAN_AT,
134
+ fixture_sha=fixture_sha,
135
+ metrics=metrics,
136
+ per_item=_collect_per_item(
137
+ inferred=inferred,
138
+ declared_fks=declared_fks,
139
+ predicted_roles=artifacts_a.roles,
140
+ gold_roles=gold_roles,
141
+ ),
142
+ notes=notes,
143
+ )
144
+
145
+ if out is None:
146
+ sys.stdout.write(json.dumps(result.to_dict(), sort_keys=True, indent=2) + "\n")
147
+ else:
148
+ write_json(result, out)
149
+
150
+
151
+ def _resolve_embeddings_flag(requested: bool, notes: list[str]) -> bool:
152
+ """Resolve the user's ``--embeddings`` request against the install state.
153
+
154
+ Returns the ``embeddings_enabled`` value the JSON envelope should
155
+ record. Appends an install-hint to ``notes`` when the user asked
156
+ for embeddings but the optional extra is not present.
157
+ """
158
+ if not requested:
159
+ return False
160
+ try:
161
+ importlib.import_module(_EMBEDDINGS_SENTINEL)
162
+ except ImportError:
163
+ notes.append(
164
+ "--embeddings was requested but the 'embeddings' extra is not "
165
+ "installed; emitted heuristic-only metrics. Install with "
166
+ "`uv sync --extra embeddings` to enable the embedding path."
167
+ )
168
+ return False
169
+ return True
170
+
171
+
172
+ def _resolve_version() -> str:
173
+ try:
174
+ return importlib.metadata.version("pretensor")
175
+ except importlib.metadata.PackageNotFoundError:
176
+ return "0.0.0+unknown"
177
+
178
+
179
+ def _load_gold_roles(dataset: Dataset) -> dict[str, str] | None:
180
+ path = _FIXTURES_ROOT / "schemas" / f"{dataset.value}_roles.yaml"
181
+ if not path.exists():
182
+ return None
183
+ raw = yaml.safe_load(path.read_text(encoding="utf-8"))
184
+ if not isinstance(raw, dict):
185
+ return None
186
+ roles = raw.get("roles") if "roles" in raw else raw
187
+ if not isinstance(roles, dict):
188
+ return None
189
+ return {str(k): str(v) for k, v in roles.items()}
190
+
191
+
192
+ def _collect_per_item(
193
+ *,
194
+ inferred: list[tuple[str, str, str, str]],
195
+ declared_fks: list[tuple[str, str, str, str]],
196
+ predicted_roles: dict[str, str],
197
+ gold_roles: dict[str, str] | None,
198
+ ) -> list[dict[str, object]]:
199
+ """Emit a deterministic per-assertion record list for the JSON output."""
200
+ items: list[dict[str, object]] = []
201
+
202
+ declared_canon = {canonicalise_join_key(k) for k in declared_fks}
203
+ inferred_canon = {canonicalise_join_key(k) for k in inferred}
204
+ edge_ids = sorted(declared_canon | inferred_canon)
205
+ for edge in edge_ids:
206
+ (src_t, src_c), (dst_t, dst_c) = edge
207
+ items.append(
208
+ {
209
+ "id": f"{src_t}.{src_c}↔{dst_t}.{dst_c}",
210
+ "kind": "inferred_join",
211
+ "expected": edge in declared_canon,
212
+ "predicted": edge in inferred_canon,
213
+ }
214
+ )
215
+
216
+ if gold_roles is not None:
217
+ for table_name in sorted(gold_roles):
218
+ # The bare table name lookup matches the format produced by
219
+ # ``build_l1_artifacts`` (``schema.table``) — gold uses bare
220
+ # names without schema, so we resolve via suffix match.
221
+ predicted = _resolve_predicted_role(predicted_roles, table_name)
222
+ items.append(
223
+ {
224
+ "id": f"role:{table_name}",
225
+ "kind": "role_classification",
226
+ "expected": gold_roles[table_name],
227
+ "predicted": predicted,
228
+ }
229
+ )
230
+
231
+ return items
232
+
233
+
234
+ def _resolve_predicted_role(
235
+ predicted: dict[str, str],
236
+ bare_table_name: str,
237
+ ) -> str | None:
238
+ """Look up a role by bare table name across the predictor's keyed-by-schema dict."""
239
+ if bare_table_name in predicted:
240
+ return predicted[bare_table_name]
241
+ for key, role in predicted.items():
242
+ # Predicted keys are ``schema.table``; gold is bare ``table``.
243
+ if key.endswith(f".{bare_table_name}"):
244
+ return role
245
+ return None
@@ -0,0 +1,27 @@
1
+ """L2 MCP-tool-quality benchmark — metric implementations and runner.
2
+
3
+ See ``docs/specs/benchmark/spec.md`` §L2 metric definitions for the
4
+ contract this module fulfils. Metric values are wrapped in the
5
+ ``{value, direction}`` envelope from
6
+ :mod:`pretensor.benchmark.results`.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ from pretensor.benchmark.l2.metrics import (
12
+ compile_metric_correctness,
13
+ query_recall_at_k,
14
+ semantic_search_recall_at_k,
15
+ top_k_with_ties,
16
+ traverse_correctness,
17
+ )
18
+ from pretensor.benchmark.l2.runner import run_l2
19
+
20
+ __all__ = [
21
+ "compile_metric_correctness",
22
+ "query_recall_at_k",
23
+ "run_l2",
24
+ "semantic_search_recall_at_k",
25
+ "top_k_with_ties",
26
+ "traverse_correctness",
27
+ ]
@@ -0,0 +1,236 @@
1
+ """Gold-data loaders for the L2 benchmark.
2
+
3
+ Reads the per-dataset NL→SQL bench JSON and the metric-templates YAML
4
+ file into typed dataclasses the runner consumes. Centralised here so
5
+ the runner stays thin and the loaders are independently testable.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+ from dataclasses import dataclass, field
12
+ from typing import Any
13
+
14
+ import yaml
15
+
16
+ from pretensor.benchmark.fixtures import Fixture
17
+ from pretensor.benchmark.l2.metrics import JoinPair
18
+ from pretensor.connectors.lineage_sqlglot import table_refs_from_sql
19
+
20
+ __all__ = [
21
+ "MetricTemplateEntry",
22
+ "QueryGoldEntry",
23
+ "TraverseGoldEntry",
24
+ "bare_table_name",
25
+ "default_schema_for",
26
+ "load_metric_templates",
27
+ "load_query_gold",
28
+ "load_traverse_gold",
29
+ ]
30
+
31
+
32
+ # Default schema applied to unqualified table references in expected_sql.
33
+ # Every gold SQL we ship today either uses public.* (Pagila, TPC-H) or is
34
+ # fully schema-qualified (AdventureWorks: humanresources.*, sales.*, etc.),
35
+ # so the default is only an absolute-last-resort fallback. Datasets we
36
+ # don't list here use ``"public"`` for the same reason.
37
+ _DEFAULT_SCHEMAS: dict[str, str] = {
38
+ "pagila": "public",
39
+ "tpch": "public",
40
+ }
41
+
42
+
43
+ @dataclass(frozen=True, slots=True)
44
+ class QueryGoldEntry:
45
+ """One question→tables observation for the ``query`` Recall@K metric."""
46
+
47
+ id: str
48
+ question: str
49
+ expected_sql: str
50
+ tables_touched: frozenset[str]
51
+
52
+
53
+ @dataclass(frozen=True, slots=True)
54
+ class TraverseGoldEntry:
55
+ """One ``(from_table, to_table, gold_path)`` entry for traverse correctness.
56
+
57
+ ``gold_path`` is the canonical sequence of ``(from_table, to_table)``
58
+ hops (schema-qualified bare names). The first hop's ``from_table`` is
59
+ the traverse start; the last hop's ``to_table`` is the destination.
60
+ """
61
+
62
+ id: str
63
+ from_table: str
64
+ to_table: str
65
+ gold_path: tuple[JoinPair, ...]
66
+
67
+
68
+ @dataclass(frozen=True, slots=True)
69
+ class MetricTemplateEntry:
70
+ """One semantic-layer metric template for ``compile_metric``."""
71
+
72
+ metric: str
73
+ semantic_yaml: str
74
+ database: str
75
+ # Optional sidecar for diagnostics — never used by the metric value.
76
+ notes: str = field(default="")
77
+
78
+
79
+ def default_schema_for(dataset_name: str) -> str:
80
+ """Return the default schema to use when SQL refs are unqualified."""
81
+ return _DEFAULT_SCHEMAS.get(dataset_name, "public")
82
+
83
+
84
+ def load_query_gold(fixture: Fixture) -> list[QueryGoldEntry]:
85
+ """Load the bench JSON and resolve ``tables_touched`` for each entry.
86
+
87
+ When an entry already has an explicit ``tables_touched`` array we use
88
+ it verbatim. When absent, we parse ``expected_sql`` via sqlglot and
89
+ extract the bare table names. Entries whose ``expected_sql`` fails to
90
+ parse are returned with an empty ``tables_touched`` set so the
91
+ metric loop skips them — the runner notes the parse failure.
92
+ """
93
+ if fixture.questions_path is None:
94
+ return []
95
+ raw = json.loads(fixture.questions_path.read_text(encoding="utf-8"))
96
+ if not isinstance(raw, list):
97
+ return []
98
+
99
+ default_schema = default_schema_for(fixture.name.value)
100
+ entries: list[QueryGoldEntry] = []
101
+ for item in raw:
102
+ if not isinstance(item, dict):
103
+ continue
104
+ identifier = str(item.get("id", "")).strip()
105
+ question = str(item.get("question", ""))
106
+ expected_sql = str(item.get("expected_sql", ""))
107
+ if not identifier:
108
+ continue
109
+
110
+ tables = item.get("tables_touched")
111
+ if isinstance(tables, list) and all(isinstance(t, str) for t in tables):
112
+ tables_set = frozenset(bare_table_name(t) for t in tables)
113
+ else:
114
+ tables_set = frozenset(
115
+ t
116
+ for _, t in table_refs_from_sql(
117
+ expected_sql,
118
+ dialect="postgres",
119
+ default_schema=default_schema,
120
+ )
121
+ )
122
+
123
+ entries.append(
124
+ QueryGoldEntry(
125
+ id=identifier,
126
+ question=question,
127
+ expected_sql=expected_sql,
128
+ tables_touched=tables_set,
129
+ )
130
+ )
131
+ return entries
132
+
133
+
134
+ def load_traverse_gold(fixture: Fixture) -> list[TraverseGoldEntry]:
135
+ """Pull ``gold_path`` entries out of the bench JSON.
136
+
137
+ Only entries that explicitly carry a ``gold_path`` array are returned
138
+ — the metric ignores entries authored only for ``query`` Recall@K.
139
+ """
140
+ if fixture.questions_path is None:
141
+ return []
142
+ raw = json.loads(fixture.questions_path.read_text(encoding="utf-8"))
143
+ if not isinstance(raw, list):
144
+ return []
145
+
146
+ out: list[TraverseGoldEntry] = []
147
+ for item in raw:
148
+ if not isinstance(item, dict):
149
+ continue
150
+ gold_path_raw = item.get("gold_path")
151
+ if not isinstance(gold_path_raw, list) or not gold_path_raw:
152
+ continue
153
+
154
+ hops: list[JoinPair] = []
155
+ for hop in gold_path_raw:
156
+ if not isinstance(hop, dict):
157
+ continue
158
+ frm = hop.get("from_table")
159
+ to = hop.get("to_table")
160
+ if isinstance(frm, str) and isinstance(to, str):
161
+ hops.append((frm, to))
162
+ if not hops:
163
+ continue
164
+
165
+ identifier = str(item.get("id", "")).strip()
166
+ out.append(
167
+ TraverseGoldEntry(
168
+ id=identifier or f"{hops[0][0]}->{hops[-1][1]}",
169
+ from_table=hops[0][0],
170
+ to_table=hops[-1][1],
171
+ gold_path=tuple(hops),
172
+ )
173
+ )
174
+ return out
175
+
176
+
177
+ def load_metric_templates(fixture: Fixture) -> list[MetricTemplateEntry]:
178
+ """Parse the per-dataset metric-templates YAML.
179
+
180
+ The YAML file lives at ``scripts/data/<dataset>_metric_templates.yaml``
181
+ and shadows :class:`Fixture.metric_templates_path`. Returns an empty
182
+ list when the file is absent — the metric reports
183
+ :class:`compile_metric_correctness` as ``None`` upstream.
184
+
185
+ Format:
186
+
187
+ .. code-block:: yaml
188
+
189
+ connection_name: pagila
190
+ templates:
191
+ - metric: total_rentals
192
+ notes: Optional human-readable rationale
193
+ semantic_yaml: |
194
+ connection_name: pagila
195
+ domains:
196
+ - name: rentals
197
+ ...
198
+
199
+ The ``connection_name`` at the top level is forwarded as the
200
+ ``database`` argument to ``compile_metric_payload``.
201
+ """
202
+ path = fixture.metric_templates_path
203
+ if path is None or not path.exists():
204
+ return []
205
+
206
+ raw: Any = yaml.safe_load(path.read_text(encoding="utf-8"))
207
+ if not isinstance(raw, dict):
208
+ return []
209
+
210
+ database = str(raw.get("connection_name", fixture.name.value))
211
+ templates_raw = raw.get("templates", [])
212
+ if not isinstance(templates_raw, list):
213
+ return []
214
+
215
+ out: list[MetricTemplateEntry] = []
216
+ for item in templates_raw:
217
+ if not isinstance(item, dict):
218
+ continue
219
+ metric = str(item.get("metric", "")).strip()
220
+ semantic_yaml = str(item.get("semantic_yaml", ""))
221
+ if not metric or not semantic_yaml.strip():
222
+ continue
223
+ out.append(
224
+ MetricTemplateEntry(
225
+ metric=metric,
226
+ semantic_yaml=semantic_yaml,
227
+ database=database,
228
+ notes=str(item.get("notes", "")),
229
+ )
230
+ )
231
+ return out
232
+
233
+
234
+ def bare_table_name(name: str) -> str:
235
+ """Strip the schema prefix from a possibly-qualified table name."""
236
+ return name.split(".", 1)[1] if "." in name else name