"""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, site_timezone: str | None = None ) -> dict[str, Any]: """Cap a Cube query, require an explicit time window, and pin the timezone. 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. An unpinned TIMEZONE is worse, because it does not look unpinned. Cube defaults to UTC, while rolling_window() builds its boundary strings in SITE_TIMEZONE - so a query without a timezone silently runs over a window shifted by the site's UTC offset while the answer states the local one. Ten hours, at this site. It is set here, at the one point every Cube query passes through, rather than in each query builder, because "remember to add the timezone" is not a control. """ capped = dict(query) limit = capped.get("limit") if not isinstance(limit, int) or limit > max_rows: capped["limit"] = max_rows if site_timezone and not capped.get("timezone"): capped["timezone"] = site_timezone 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")