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,316 @@
1
+ """LLM tool-use agent loop for the L3 pretensor runner.
2
+
3
+ The runner asks the LLM to translate one NL question into PostgreSQL,
4
+ giving the model the MCP tool set discovered from ``pretensor serve``.
5
+ The loop here is provider-agnostic: it owns the message history, calls
6
+ the LLM (via :class:`AgentLlmClient`), routes tool invocations through
7
+ a caller-supplied callback, and stops as soon as the model returns
8
+ final text (the SQL).
9
+
10
+ A hard cap on iterations keeps a confused model from running away
11
+ forever. Per-call telemetry (latency, tokens, tool trace) is returned
12
+ so the runner can populate the JSON envelope without a second pass.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ from dataclasses import dataclass, field
18
+ from typing import Any, Callable, Literal, Protocol
19
+
20
+ from pretensor.errors import PretensorError
21
+
22
+ __all__ = [
23
+ "AgentLlmClient",
24
+ "AgentLoopError",
25
+ "AgentLoopResult",
26
+ "AgentMessage",
27
+ "AgentStep",
28
+ "AgentTool",
29
+ "AgentToolCall",
30
+ "AgentToolResult",
31
+ "DEFAULT_MAX_ITERATIONS",
32
+ "ToolCallTraceEntry",
33
+ "ToolInvocationOutcome",
34
+ "ToolInvoker",
35
+ "run_agent_loop",
36
+ ]
37
+
38
+
39
+ DEFAULT_MAX_ITERATIONS = 16
40
+ """Cap on the LLM ↔ tool round-trips per question.
41
+
42
+ Pagila's hardest gold question solves in ≤4 round-trips today; 16
43
+ leaves headroom for messier schemas without letting a confused model
44
+ spin forever. Hitting the cap is recorded in ``notes[]`` and the
45
+ question is marked failed — never silently accepted.
46
+ """
47
+
48
+
49
+ _TOOL_RESULT_TRUNCATE_BYTES = 8000
50
+ """Soft cap on the tool result handed back to the LLM, in UTF-8 bytes.
51
+
52
+ The MCP server can return large payloads (e.g. ``cypher`` over a
53
+ big graph). Trimming protects the next request's prompt budget; the
54
+ tool trace stores the *full* ``response_size`` (also in bytes) so
55
+ an auditor can see when truncation kicked in. Both values share the
56
+ same unit so comparing them is unambiguous.
57
+ """
58
+
59
+ _TRUNCATION_SUFFIX = "…[truncated]"
60
+
61
+
62
+ @dataclass(frozen=True, slots=True)
63
+ class AgentTool:
64
+ """An MCP tool exposed to the agent.
65
+
66
+ Mirrors :class:`mcp.types.Tool` but lives in the L3 namespace so
67
+ callers don't have to import the MCP package to read the trace.
68
+ """
69
+
70
+ name: str
71
+ description: str
72
+ input_schema: dict[str, Any]
73
+
74
+
75
+ @dataclass(frozen=True, slots=True)
76
+ class AgentToolCall:
77
+ """One tool invocation requested by the LLM.
78
+
79
+ ``id`` is the provider's identifier (Anthropic ``tool_use_id``,
80
+ OpenAI ``tool_call_id``); we forward it back unchanged so the
81
+ next request can pair tool results with their calls.
82
+ """
83
+
84
+ id: str
85
+ name: str
86
+ arguments: dict[str, Any]
87
+
88
+
89
+ @dataclass(frozen=True, slots=True)
90
+ class AgentToolResult:
91
+ """Result of executing one :class:`AgentToolCall`."""
92
+
93
+ id: str
94
+ content: str
95
+ is_error: bool = False
96
+
97
+
98
+ @dataclass(frozen=True, slots=True)
99
+ class AgentMessage:
100
+ """One turn of the LLM conversation, provider-agnostic.
101
+
102
+ Role is either ``"assistant"`` (LLM produced this turn) or
103
+ ``"user"`` (we produced this turn — initial question, or a batch
104
+ of tool results). A single assistant turn may contain text and
105
+ one or more tool calls; a single user turn either holds the
106
+ initial prompt text or a batch of tool results.
107
+ """
108
+
109
+ role: Literal["user", "assistant"]
110
+ text: str | None = None
111
+ tool_calls: tuple[AgentToolCall, ...] = ()
112
+ tool_results: tuple[AgentToolResult, ...] = ()
113
+
114
+
115
+ @dataclass(frozen=True, slots=True)
116
+ class AgentStep:
117
+ """One round-trip with the LLM.
118
+
119
+ Either ``text`` is set (the model finished and returned a final
120
+ answer — for L3, the SQL) or ``tool_calls`` is non-empty (the
121
+ model wants to call tools before producing text). Token counts
122
+ are optional; some providers omit them.
123
+ """
124
+
125
+ text: str | None
126
+ tool_calls: tuple[AgentToolCall, ...]
127
+ prompt_tokens: int | None
128
+ completion_tokens: int | None
129
+ latency_ms: int
130
+
131
+
132
+ class AgentLlmClient(Protocol):
133
+ """Single-step tool-use surface for the L3 pretensor runner.
134
+
135
+ Implementations translate :class:`AgentMessage` history to the
136
+ provider's wire format (Anthropic content blocks, OpenAI
137
+ ``tool_calls``, etc.) and back. The agent loop does not see the
138
+ provider-specific shape.
139
+ """
140
+
141
+ def agent_complete(
142
+ self,
143
+ *,
144
+ system: str,
145
+ messages: list[AgentMessage],
146
+ tools: list[AgentTool],
147
+ model: str,
148
+ temperature: float,
149
+ ) -> AgentStep: ...
150
+
151
+
152
+ class AgentLoopError(PretensorError, RuntimeError):
153
+ """Raised when the loop cannot make further progress.
154
+
155
+ Specifically: the iteration cap was hit without the LLM returning
156
+ final text, OR the LLM produced an empty step (no text and no
157
+ tool calls — providers shouldn't but it's not impossible).
158
+ """
159
+
160
+
161
+ @dataclass(frozen=True, slots=True)
162
+ class ToolCallTraceEntry:
163
+ """One row of the per-question tool-call trace.
164
+
165
+ Mirrors AC-mandated shape: ``{tool, args, response_size}``.
166
+ ``is_error`` flags MCP-level failures (subprocess error, unknown
167
+ tool, handler exception); the runner uses it to surface failed
168
+ runs without losing the rest of the trace.
169
+ """
170
+
171
+ tool: str
172
+ args: dict[str, Any]
173
+ response_size: int
174
+ is_error: bool
175
+
176
+
177
+ @dataclass(slots=True)
178
+ class AgentLoopResult:
179
+ """Final outcome of the loop for one question.
180
+
181
+ ``text`` is the LLM's last textual reply (the SQL, with no fence
182
+ cleanup yet — the runner does that). ``trace`` is the ordered list
183
+ of tool calls; ``iterations`` is how many LLM round-trips ran. Token
184
+ and latency totals are sums across all steps.
185
+ """
186
+
187
+ text: str
188
+ trace: list[ToolCallTraceEntry] = field(default_factory=list)
189
+ iterations: int = 0
190
+ total_prompt_tokens: int | None = None
191
+ total_completion_tokens: int | None = None
192
+ total_llm_latency_ms: int = 0
193
+
194
+
195
+ @dataclass(frozen=True, slots=True)
196
+ class ToolInvocationOutcome:
197
+ """Outcome of one tool invocation, ready to attach back to the LLM.
198
+
199
+ Plain ``content`` + ``is_error`` so the runner doesn't have to know
200
+ about provider-specific tool_use ids — the loop pairs the outcome
201
+ with the originating call when it builds the next user turn.
202
+ """
203
+
204
+ content: str
205
+ is_error: bool = False
206
+
207
+
208
+ ToolInvoker = Callable[[str, dict[str, Any]], ToolInvocationOutcome]
209
+ """Sync callable the loop uses to run one tool.
210
+
211
+ The runner provides this — typically wrapping an
212
+ :class:`McpClient.call_tool` plus serialisation. Returning a
213
+ :class:`ToolInvocationOutcome` (rather than a raw dict) lets the runner
214
+ mark tool-level errors that the LLM should still see.
215
+ """
216
+
217
+
218
+ def run_agent_loop(
219
+ *,
220
+ client: AgentLlmClient,
221
+ system: str,
222
+ user: str,
223
+ tools: list[AgentTool],
224
+ invoke_tool: ToolInvoker,
225
+ model: str,
226
+ temperature: float,
227
+ max_iterations: int = DEFAULT_MAX_ITERATIONS,
228
+ ) -> AgentLoopResult:
229
+ """Run one question to completion against ``client`` and ``invoke_tool``.
230
+
231
+ Each iteration sends the full conversation history to the LLM. If
232
+ the model returns text, the loop exits with that text. If the model
233
+ returns tool calls, each call is executed (via ``invoke_tool``),
234
+ the results are appended as a user turn, and the loop continues.
235
+ """
236
+ messages: list[AgentMessage] = [AgentMessage(role="user", text=user)]
237
+ trace: list[ToolCallTraceEntry] = []
238
+ total_prompt_tokens: int | None = None
239
+ total_completion_tokens: int | None = None
240
+ total_latency_ms = 0
241
+
242
+ for iteration in range(1, max_iterations + 1):
243
+ step = client.agent_complete(
244
+ system=system,
245
+ messages=messages,
246
+ tools=tools,
247
+ model=model,
248
+ temperature=temperature,
249
+ )
250
+ total_latency_ms += step.latency_ms
251
+ if step.prompt_tokens is not None:
252
+ total_prompt_tokens = (total_prompt_tokens or 0) + step.prompt_tokens
253
+ if step.completion_tokens is not None:
254
+ total_completion_tokens = (
255
+ total_completion_tokens or 0
256
+ ) + step.completion_tokens
257
+
258
+ if step.text is not None and not step.tool_calls:
259
+ return AgentLoopResult(
260
+ text=step.text,
261
+ trace=trace,
262
+ iterations=iteration,
263
+ total_prompt_tokens=total_prompt_tokens,
264
+ total_completion_tokens=total_completion_tokens,
265
+ total_llm_latency_ms=total_latency_ms,
266
+ )
267
+
268
+ if not step.tool_calls:
269
+ raise AgentLoopError(
270
+ f"LLM returned an empty step at iteration {iteration} "
271
+ "(no text and no tool calls)."
272
+ )
273
+
274
+ # Record the assistant turn (text + tool calls) so the next request
275
+ # carries the full history. ``step.text`` may be non-empty when the
276
+ # model "thinks out loud" before calling a tool — keep it.
277
+ messages.append(
278
+ AgentMessage(
279
+ role="assistant",
280
+ text=step.text,
281
+ tool_calls=step.tool_calls,
282
+ )
283
+ )
284
+
285
+ results: list[AgentToolResult] = []
286
+ for call in step.tool_calls:
287
+ outcome = invoke_tool(call.name, call.arguments)
288
+ content_bytes = outcome.content.encode("utf-8")
289
+ trace.append(
290
+ ToolCallTraceEntry(
291
+ tool=call.name,
292
+ args=dict(call.arguments),
293
+ response_size=len(content_bytes),
294
+ is_error=outcome.is_error,
295
+ )
296
+ )
297
+ content = outcome.content
298
+ if len(content_bytes) > _TOOL_RESULT_TRUNCATE_BYTES:
299
+ # Slice on bytes to keep the truncation cap unambiguous;
300
+ # decode with ``errors="ignore"`` so a multibyte
301
+ # codepoint straddling the boundary doesn't raise.
302
+ content = (
303
+ content_bytes[:_TOOL_RESULT_TRUNCATE_BYTES].decode(
304
+ "utf-8", errors="ignore"
305
+ )
306
+ + _TRUNCATION_SUFFIX
307
+ )
308
+ results.append(
309
+ AgentToolResult(id=call.id, content=content, is_error=outcome.is_error)
310
+ )
311
+
312
+ messages.append(AgentMessage(role="user", tool_results=tuple(results)))
313
+
314
+ raise AgentLoopError(
315
+ f"Agent did not produce final text within {max_iterations} iterations."
316
+ )
@@ -0,0 +1,188 @@
1
+ """Read-only SQL execution helpers for the L3 benchmark runners.
2
+
3
+ The runners need to execute both the gold SQL and the agent-emitted SQL
4
+ against a real PostgreSQL instance, then compare the resulting rows. This
5
+ module is the only place that opens DB connections in the L3 layer, so all
6
+ the safety guardrails live here:
7
+
8
+ * **Database URL is resolved from a per-dataset env var.** No CLI flag —
9
+ ``--dataset`` already pins which DB to talk to, and growing the CLI
10
+ surface fragments the contract.
11
+ * **Every query runs in a transaction that is rolled back at the end.**
12
+ Defense-in-depth against a misbehaving agent that emits ``DROP TABLE``;
13
+ the read-only intent of the run is preserved even if the syntactic
14
+ guard below misses a corner case.
15
+ * **Statement-level timeout** (30s by default) prevents pathological
16
+ agent SQL (cross joins, accidental cartesian products) from hanging
17
+ the runner.
18
+ * **Non-SELECT guard** rejects any agent SQL whose first non-comment,
19
+ non-whitespace token is not ``SELECT``, ``WITH``, or ``(``. Together
20
+ with the rolled-back transaction this is two independent layers of
21
+ protection — a slip in either is contained by the other.
22
+ """
23
+
24
+ from __future__ import annotations
25
+
26
+ import os
27
+ import re
28
+ from typing import Any
29
+
30
+ import sqlalchemy
31
+ from sqlalchemy.engine import Engine
32
+ from sqlalchemy.exc import SQLAlchemyError
33
+
34
+ from pretensor.benchmark.runner import Dataset
35
+ from pretensor.errors import PretensorError
36
+
37
+ __all__ = [
38
+ "DEFAULT_STATEMENT_TIMEOUT_MS",
39
+ "QueryExecutionError",
40
+ "RefusedNonSelectError",
41
+ "execute_query",
42
+ "is_select_only",
43
+ "resolve_database_url",
44
+ ]
45
+
46
+
47
+ DEFAULT_STATEMENT_TIMEOUT_MS = 30_000
48
+
49
+
50
+ _DATASET_TO_ENV_VAR: dict[Dataset, str] = {
51
+ Dataset.PAGILA: "PAGILA_DATABASE_URL",
52
+ Dataset.TPCH: "TPCH_DATABASE_URL",
53
+ Dataset.ADVENTUREWORKS: "ADVENTUREWORKS_DATABASE_URL",
54
+ }
55
+ """L3 baseline supports the three OSS-bundled datasets that ship with DDL.
56
+
57
+ The synthetic fixtures (``analytics_dwh``, ``saas_multitenant``,
58
+ ``adversarial``) have no DDL dump and no live DB they correspond to,
59
+ so they are deliberately absent here — calling ``resolve_database_url``
60
+ on them raises a clear error.
61
+ """
62
+
63
+
64
+ class QueryExecutionError(PretensorError, RuntimeError):
65
+ """Raised when a SQL statement fails to execute (parse error, missing
66
+ object, runtime fault, statement timeout, etc.)."""
67
+
68
+
69
+ class RefusedNonSelectError(QueryExecutionError):
70
+ """Raised when an agent SQL statement is rejected by the SELECT-only guard.
71
+
72
+ A subclass of ``QueryExecutionError`` so callers that want to record
73
+ "could not execute" can use a single ``except`` and still inspect the
74
+ type to distinguish "refused before execution" from "DB error".
75
+ """
76
+
77
+
78
+ def resolve_database_url(dataset: Dataset, env: dict[str, str] | None = None) -> str:
79
+ """Read the per-dataset DSN env var; raise if unset or dataset unsupported.
80
+
81
+ ``env`` is for tests (pass an explicit dict to bypass ``os.environ``);
82
+ production callers leave it ``None`` and the function reads the
83
+ process environment.
84
+ """
85
+ if env is None:
86
+ env = dict(os.environ)
87
+ try:
88
+ var_name = _DATASET_TO_ENV_VAR[dataset]
89
+ except KeyError as exc:
90
+ raise LookupError(
91
+ f"L3 baseline runner does not support dataset {dataset.value!r}: "
92
+ f"only {[d.value for d in _DATASET_TO_ENV_VAR]} have DDL bundles."
93
+ ) from exc
94
+ value = env.get(var_name)
95
+ if not value:
96
+ raise LookupError(
97
+ f"Set {var_name} to the PostgreSQL DSN of a {dataset.value} "
98
+ f"database (e.g. {var_name}=postgresql://user@host:5432/{dataset.value})."
99
+ )
100
+ return value
101
+
102
+
103
+ # Strip line comments (``--`` to end of line) and block comments (``/* ... */``);
104
+ # whitespace; then peek at the leading token. CTEs (``WITH ...``) and
105
+ # parenthesised SELECTs (``(SELECT ...) UNION ...``) are allowed alongside
106
+ # bare ``SELECT``.
107
+ _LINE_COMMENT_RE = re.compile(r"--[^\n]*")
108
+ _BLOCK_COMMENT_RE = re.compile(r"/\*.*?\*/", re.DOTALL)
109
+ _ALLOWED_LEAD_TOKENS = ("SELECT", "WITH", "TABLE", "VALUES", "(")
110
+
111
+
112
+ def is_select_only(sql: str) -> bool:
113
+ """Return ``True`` iff ``sql`` begins with a read-only top-level keyword.
114
+
115
+ ``TABLE`` and ``VALUES`` are PostgreSQL read forms; ``WITH`` covers
116
+ CTEs whose final statement is a ``SELECT`` (we don't deep-parse the
117
+ CTE — combined with the rolled-back transaction in
118
+ :func:`execute_query`, that's enough). A bare ``(`` allows
119
+ ``(SELECT ...) UNION ...`` style queries.
120
+
121
+ A statement that *contains* DDL/DML keywords later (in a string
122
+ literal, in a column comment, etc.) is not blocked here — that's a
123
+ job for the database's own parser. The only goal of this guard is
124
+ "first thing the LLM emits must be a read".
125
+ """
126
+ stripped = _BLOCK_COMMENT_RE.sub(" ", sql)
127
+ stripped = _LINE_COMMENT_RE.sub(" ", stripped).strip()
128
+ if not stripped:
129
+ return False
130
+ upper = stripped.upper()
131
+ for tok in _ALLOWED_LEAD_TOKENS:
132
+ if tok == "(" and upper.startswith("("):
133
+ return True
134
+ if upper.startswith(tok) and (
135
+ len(upper) == len(tok) or not upper[len(tok)].isalnum()
136
+ ):
137
+ return True
138
+ return False
139
+
140
+
141
+ def execute_query(
142
+ dsn: str,
143
+ sql: str,
144
+ *,
145
+ enforce_select_only: bool = False,
146
+ statement_timeout_ms: int = DEFAULT_STATEMENT_TIMEOUT_MS,
147
+ ) -> tuple[list[tuple[Any, ...]], list[str]]:
148
+ """Run ``sql`` against ``dsn`` inside a rolled-back transaction.
149
+
150
+ Returns ``(rows, column_names)``. ``enforce_select_only=True`` runs
151
+ the syntactic guard first and raises :class:`RefusedNonSelectError`
152
+ when it fails — used for agent-emitted SQL. Gold SQL skips the guard
153
+ (it's authored by us and may use any read form, but we still wrap
154
+ in a rolled-back transaction).
155
+
156
+ Any DB error (parse failure, missing object, statement timeout, etc.)
157
+ is wrapped in :class:`QueryExecutionError` so callers can record it
158
+ without having to know the SQLAlchemy / driver-specific exception
159
+ hierarchy.
160
+ """
161
+ if enforce_select_only and not is_select_only(sql):
162
+ raise RefusedNonSelectError(
163
+ "agent SQL refused: first non-comment token is not SELECT/WITH/(/VALUES."
164
+ )
165
+ # A new Engine per call is intentional: L3 only runs ~10–25 questions
166
+ # per dataset (so 2N ≤ ~50 engines), each query owns its own pool /
167
+ # transaction lifecycle, and the dispose() in the finally block keeps
168
+ # connection state from leaking across questions if anything goes wrong.
169
+ engine: Engine = sqlalchemy.create_engine(dsn)
170
+ try:
171
+ with engine.connect() as conn:
172
+ trans = conn.begin()
173
+ try:
174
+ conn.execute(
175
+ sqlalchemy.text(
176
+ f"SET LOCAL statement_timeout = {int(statement_timeout_ms)}"
177
+ )
178
+ )
179
+ cursor = conn.execute(sqlalchemy.text(sql))
180
+ rows = [tuple(row) for row in cursor.fetchall()]
181
+ column_names = list(cursor.keys())
182
+ return rows, column_names
183
+ except SQLAlchemyError as exc:
184
+ raise QueryExecutionError(str(exc)) from exc
185
+ finally:
186
+ trans.rollback()
187
+ finally:
188
+ engine.dispose()
@@ -0,0 +1,85 @@
1
+ """L3 gold-question loader.
2
+
3
+ The L3 gold corpus lives at ``scripts/data/<dataset>_nl2sql_bench.json`` —
4
+ the same file L2's ``query_recall`` reads. The L3 baseline runner only
5
+ needs three fields per question (``id``, ``question``, ``expected_sql``),
6
+ so this loader strips the rest and returns a sorted list for byte-stable
7
+ output ordering.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import json
13
+ from dataclasses import dataclass
14
+ from pathlib import Path
15
+
16
+ from pretensor.benchmark.fixtures import Fixture
17
+
18
+ __all__ = ["L3GoldEntry", "load_l3_gold"]
19
+
20
+
21
+ @dataclass(frozen=True, slots=True)
22
+ class L3GoldEntry:
23
+ """One NL-to-SQL question to grade.
24
+
25
+ ``id`` is the stable identifier used in the JSON envelope's
26
+ ``per_item`` list. ``question`` is the natural-language prompt the
27
+ runner sends to the LLM. ``expected_sql`` is the human-authored gold
28
+ query whose result rows define the correct answer.
29
+ """
30
+
31
+ id: str
32
+ question: str
33
+ expected_sql: str
34
+
35
+
36
+ def load_l3_gold(fixture: Fixture) -> tuple[Path, bytes, list[L3GoldEntry]]:
37
+ """Read and parse the dataset's NL-to-SQL gold file.
38
+
39
+ Returns ``(questions_path, raw_bytes, entries)``. ``raw_bytes`` is
40
+ the file contents exactly as read from disk — exposed so callers can
41
+ fingerprint the fixture (e.g. ``hashlib.sha256``) without re-reading
42
+ the file from a second I/O. ``questions_path`` makes
43
+ ``fixture.questions_path``'s resolved value part of the function's
44
+ contract so callers don't need ``Optional`` narrowing.
45
+
46
+ Entries are sorted by ``id`` so the runner's per-item output stays
47
+ byte-stable across re-runs even if the source JSON's array order
48
+ shifts. Raises :class:`FileNotFoundError` if the dataset has no
49
+ question set on disk; raises :class:`ValueError` for malformed
50
+ records.
51
+ """
52
+ if fixture.questions_path is None:
53
+ raise FileNotFoundError(
54
+ f"Dataset {fixture.name.value!r} has no NL-to-SQL gold "
55
+ "questions file checked in (looked for "
56
+ f"scripts/data/{fixture.name.value}_nl2sql_bench.json)."
57
+ )
58
+ questions_path = fixture.questions_path
59
+ raw_bytes = questions_path.read_bytes()
60
+ raw = json.loads(raw_bytes.decode("utf-8"))
61
+ if not isinstance(raw, list):
62
+ raise ValueError(
63
+ f"{questions_path}: expected a JSON array, got {type(raw).__name__}."
64
+ )
65
+ entries: list[L3GoldEntry] = []
66
+ for index, record in enumerate(raw):
67
+ if not isinstance(record, dict):
68
+ raise ValueError(
69
+ f"{questions_path}[{index}]: expected an object, got "
70
+ f"{type(record).__name__}."
71
+ )
72
+ try:
73
+ entries.append(
74
+ L3GoldEntry(
75
+ id=str(record["id"]),
76
+ question=str(record["question"]),
77
+ expected_sql=str(record["expected_sql"]),
78
+ )
79
+ )
80
+ except KeyError as exc:
81
+ raise ValueError(
82
+ f"{questions_path}[{index}]: missing field {exc.args[0]!r}."
83
+ ) from exc
84
+ entries.sort(key=lambda e: e.id)
85
+ return questions_path, raw_bytes, entries