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,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
|