Files
ThothII/harness/tht/execute/__init__.py
T

133 lines
4.7 KiB
Python

import json
import time
from dataclasses import dataclass
import sqlglot
from sqlalchemy import Engine, text
from sqlalchemy.exc import DBAPIError
class ExecutionError(Exception):
"""Errore di esecuzione con messaggio leggibile per il reviewer."""
@dataclass
class ExecResult:
columns: list[str]
rows: list[tuple]
execution_ms: int
truncated: bool
@dataclass
class PlanSummary:
total_cost: float
plan_rows: int
node_types: list[str]
def require_positive_int(value: object, *, name: str) -> int:
"""Return a validated positive integer, excluding booleans and numeric lookalikes."""
if type(value) is not int or value <= 0:
raise ValueError(f"{name} must be a positive integer")
return value
def _inject_limit(sql: str, limit: int) -> tuple[str, bool]:
"""Aggiunge LIMIT limit+1 se assente (il +1 serve a rilevare il troncamento).
Se la query ha gia' un suo LIMIT, lo si rispetta."""
ast = sqlglot.parse_one(sql, read="postgres")
# solo le query (SELECT/UNION & co.) accettano un LIMIT: per qualunque altro
# statement non si inietta nulla e la transazione READ ONLY lo rifiutera' (rete 2).
if not isinstance(ast, sqlglot.exp.Query) or ast.args.get("limit") is not None:
return sql, False
return ast.limit(limit + 1).sql(dialect="postgres"), True
def assert_read_only(sql: str) -> None:
"""Guard read-only strutturale condiviso dai due codepath di esecuzione (D7).
Difesa in profondita': oltre alla rete server (utente RO / READ ONLY tx / RPC
SELECT-only), ogni esecuzione passa da qui, cosi' anche un riuso diretto delle API
Python (es. dal backend) non puo' bypassare il single-statement + SELECT-only.
Usa sqlcheck.validate_sql senza schema fisico: parse + un solo statement +
read-only per struttura + funzioni vietate.
"""
from tht.sqlcheck import validate_sql
check = validate_sql(sql)
if not check.ok:
raise ExecutionError(
"SQL rifiutato (read-only enforcement): " + "; ".join(check.errors)
)
def _translate_error(e: DBAPIError) -> ExecutionError:
msg = str(e.orig)
if "statement timeout" in msg or "canceling statement" in msg:
return ExecutionError("timeout: la query ha superato statement_timeout_ms")
if "read-only" in msg:
return ExecutionError("rifiutato: statement non eseguibile in transazione read-only")
return ExecutionError(f"errore SQL: {msg}")
def run_controlled(engine: Engine, sql: str, *, limit: int, timeout_ms: int) -> ExecResult:
"""Esecuzione nelle 4 reti: utente RO (a monte), transazione READ ONLY,
statement_timeout, LIMIT iniettato via AST. Guard read-only client-side a monte."""
assert_read_only(sql)
final_sql, injected = _inject_limit(sql, limit)
start = time.monotonic()
with engine.connect() as conn:
trans = conn.begin()
try:
conn.execute(text("SET TRANSACTION READ ONLY"))
conn.execute(text(f"SET LOCAL statement_timeout = {int(timeout_ms)}"))
result = conn.execute(text(final_sql))
columns = list(result.keys())
rows = [tuple(r) for r in result.fetchall()]
except DBAPIError as e:
raise _translate_error(e) from e
finally:
trans.rollback()
elapsed_ms = int((time.monotonic() - start) * 1000)
truncated = injected and len(rows) > limit
return ExecResult(
columns=columns, rows=rows[:limit], execution_ms=elapsed_ms, truncated=truncated
)
def explain(engine: Engine, sql: str, *, timeout_ms: int) -> PlanSummary:
"""EXPLAIN (FORMAT JSON), mai ANALYZE: il piano si stima, non si esegue.
assert_read_only a monte: l'EXPLAIN interpola lo SQL, quindi il single-statement +
SELECT-only va garantito anche qui (non solo nel chiamante CLI)."""
assert_read_only(sql)
with engine.connect() as conn:
trans = conn.begin()
try:
conn.execute(text("SET TRANSACTION READ ONLY"))
conn.execute(text(f"SET LOCAL statement_timeout = {int(timeout_ms)}"))
row = conn.execute(text(f"EXPLAIN (FORMAT JSON) {sql}")).fetchone()
except DBAPIError as e:
raise _translate_error(e) from e
finally:
trans.rollback()
payload = row[0]
if isinstance(payload, str):
payload = json.loads(payload)
plan = payload[0]["Plan"]
node_types: list[str] = []
def collect(node: dict) -> None:
node_types.append(node.get("Node Type", "?"))
for child in node.get("Plans", []):
collect(child)
collect(plan)
return PlanSummary(
total_cost=float(plan.get("Total Cost", 0.0)),
plan_rows=int(plan.get("Plan Rows", 0)),
node_types=node_types,
)