yau-plant-assistant/api/guardrails.py
Claude 34d2ccc576 Scaffold the WRPS plant operations assistant repository
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>
2026-08-20 13:56:32 +10:00

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")