Build spec and host brief carried in from C:\Claude and WRPS/02-env; the plant model (equipment, tags, alarm bitmask, enums, unit conversions) is derived from WRPS/04-plc/register-map.csv, WRPS/05-scada/modbus/scada-points.csv and WRPS-CTL-003. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
221 lines
7.7 KiB
Python
221 lines
7.7 KiB
Python
"""SQL allow-list, resource caps, and contract enforcement.
|
|
|
|
Three jobs, in order of how much they matter:
|
|
|
|
1. Nothing but a single bounded SELECT reaches a database. Enforced with
|
|
sqlglot on the parsed tree, not a regex over the string - a regex over SQL
|
|
is a suggestion.
|
|
2. Every query is capped: row limit and statement timeout.
|
|
3. Every contract rejection is logged to Langfuse with the offending output,
|
|
so the failure is visible rather than silently regenerated away.
|
|
|
|
The database role is the FIRST line of defence (agent_ro holds SELECT and
|
|
nothing else, see db/003_roles.sql). This module is the second. Neither is
|
|
sufficient alone: the role stops writes, this stops a SELECT that scans imh for
|
|
a year, and only the role stops a bug here from becoming a write.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from dataclasses import dataclass, field
|
|
from typing import Any
|
|
|
|
import sqlglot
|
|
from sqlglot import exp
|
|
|
|
from contracts import ContractViolation, QuestionClass
|
|
|
|
log = logging.getLogger("guardrails")
|
|
|
|
|
|
class GuardrailViolation(Exception):
|
|
"""A query was refused before execution."""
|
|
|
|
def __init__(self, rule: str, detail: str) -> None:
|
|
super().__init__(f"{rule}: {detail}")
|
|
self.rule = rule
|
|
self.detail = detail
|
|
|
|
|
|
# Tables the agent may read. Anything else is refused, including tables that
|
|
# exist and would be harmless - an allow-list that grows silently is not one.
|
|
ALLOWED_TABLES: set[str] = {
|
|
"equipment",
|
|
"tags",
|
|
"doc_chunks",
|
|
"fixture.alarm_history",
|
|
"fixture.process_value_history",
|
|
"fixture.operation_history",
|
|
}
|
|
|
|
# Built from names so a sqlglot upgrade that renames or removes a node type
|
|
# fails loudly at import rather than silently dropping a check.
|
|
_FORBIDDEN_NAMES = [
|
|
"Insert", "Update", "Delete", "Drop", "Create", "Alter", "Merge",
|
|
"Command", "Copy", "Grant",
|
|
]
|
|
_FORBIDDEN = tuple(
|
|
node for node in (getattr(exp, name, None) for name in _FORBIDDEN_NAMES) if node
|
|
)
|
|
|
|
# A top-level node that is a legitimate read. exp.Command covers anything
|
|
# sqlglot could not classify, and it is in the forbidden list above.
|
|
_READ_NODES = tuple(
|
|
node
|
|
for node in (getattr(exp, name, None) for name in ("Select", "Union", "With"))
|
|
if node
|
|
)
|
|
|
|
|
|
def check_sql(sql: str, *, max_rows: int, dialect: str = "postgres") -> str:
|
|
"""Parse, validate and return the SQL to execute, with a LIMIT applied.
|
|
|
|
Refuses anything that is not exactly one SELECT over allow-listed tables.
|
|
"""
|
|
try:
|
|
statements = sqlglot.parse(sql, dialect=dialect)
|
|
except Exception as exc:
|
|
raise GuardrailViolation("unparseable", str(exc)) from exc
|
|
|
|
statements = [s for s in statements if s is not None]
|
|
if len(statements) != 1:
|
|
raise GuardrailViolation(
|
|
"multiple_statements", f"{len(statements)} statements in one query"
|
|
)
|
|
|
|
stmt = statements[0]
|
|
if not isinstance(stmt, _READ_NODES):
|
|
raise GuardrailViolation("not_a_select", f"top level node is {type(stmt).__name__}")
|
|
|
|
for node in stmt.walk():
|
|
if isinstance(node, _FORBIDDEN):
|
|
raise GuardrailViolation("write_operation", type(node).__name__)
|
|
|
|
for table in stmt.find_all(exp.Table):
|
|
name = f"{table.db}.{table.name}" if table.db else table.name
|
|
if name.lower() not in {t.lower() for t in ALLOWED_TABLES}:
|
|
raise GuardrailViolation("table_not_allowed", name)
|
|
|
|
# Cap the rows. An explicit smaller limit is honoured; a larger one is not.
|
|
existing = stmt.args.get("limit")
|
|
if existing is None:
|
|
stmt = stmt.limit(max_rows)
|
|
else:
|
|
try:
|
|
if int(existing.expression.name) > max_rows:
|
|
stmt = stmt.limit(max_rows)
|
|
except (AttributeError, ValueError):
|
|
stmt = stmt.limit(max_rows)
|
|
|
|
return stmt.sql(dialect=dialect)
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Cube query caps. Cube generates its own SQL, so check_sql does not apply -
|
|
# what is validated instead is the query object the agent asked for.
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def check_cube_query(query: dict[str, Any], *, max_rows: int) -> dict[str, Any]:
|
|
"""Cap a Cube query and require an explicit time window.
|
|
|
|
An unpinned time window is the single most common way a data answer becomes
|
|
unreproducible: imh is live, so the same question asked twice gives two
|
|
answers and neither can be checked.
|
|
"""
|
|
capped = dict(query)
|
|
limit = capped.get("limit")
|
|
if not isinstance(limit, int) or limit > max_rows:
|
|
capped["limit"] = max_rows
|
|
|
|
time_dimensions = capped.get("timeDimensions") or []
|
|
if not time_dimensions:
|
|
raise GuardrailViolation(
|
|
"unpinned_time_window",
|
|
"every Cube query must carry an explicit timeDimensions range",
|
|
)
|
|
for td in time_dimensions:
|
|
if not td.get("dateRange"):
|
|
raise GuardrailViolation(
|
|
"unpinned_time_window", f"no dateRange on {td.get('dimension')}"
|
|
)
|
|
return capped
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Contract enforcement — generate, validate, regenerate ONCE, then error.
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
@dataclass
|
|
class EnforcementResult:
|
|
answer: Any
|
|
attempts: int
|
|
violations: list[ContractViolation] = field(default_factory=list)
|
|
|
|
|
|
def enforce_contract(generate, klass: QuestionClass, *, trace=None) -> EnforcementResult:
|
|
"""Run `generate()` until its output satisfies the class contract.
|
|
|
|
`generate(attempt, previous_violation)` returns a dict payload.
|
|
|
|
One retry. Not two, not "until it works" - a model that fails a safety
|
|
contract twice is not going to be argued into compliance, and each retry
|
|
costs a flagship call. The second failure raises, and the caller returns an
|
|
error to the operator.
|
|
"""
|
|
from contracts import validate_answer # local import keeps the cycle out
|
|
|
|
violations: list[ContractViolation] = []
|
|
previous: ContractViolation | None = None
|
|
|
|
for attempt in (1, 2):
|
|
payload = generate(attempt, previous)
|
|
try:
|
|
answer = validate_answer(payload, klass)
|
|
return EnforcementResult(answer=answer, attempts=attempt, violations=violations)
|
|
except ContractViolation as violation:
|
|
violations.append(violation)
|
|
previous = violation
|
|
log_violation(violation, klass, attempt, trace=trace)
|
|
|
|
raise violations[-1]
|
|
|
|
|
|
def log_violation(
|
|
violation: ContractViolation,
|
|
klass: QuestionClass,
|
|
attempt: int,
|
|
*,
|
|
trace=None,
|
|
) -> None:
|
|
"""Every rejection goes to Langfuse WITH the offending output.
|
|
|
|
The offending output is the whole value of the log line - a count of
|
|
violations tells you nothing about what the model tried to say. It stays
|
|
inside Langfuse, which is behind Authelia; it never reaches the operator
|
|
and never goes in an error message.
|
|
"""
|
|
log.warning(
|
|
"contract_violation class=%s attempt=%s rule=%s detail=%s",
|
|
klass.value,
|
|
attempt,
|
|
violation.rule,
|
|
violation.detail,
|
|
)
|
|
if trace is not None:
|
|
try:
|
|
trace.event(
|
|
name="contract_violation",
|
|
level="WARNING",
|
|
metadata={
|
|
"question_class": klass.value,
|
|
"attempt": attempt,
|
|
"rule": violation.rule,
|
|
"detail": violation.detail,
|
|
},
|
|
input=violation.offending_output,
|
|
)
|
|
except Exception: # tracing must never break the request path
|
|
log.exception("failed to record contract violation in Langfuse")
|