Files
ThothII/backend/scripts/revision_state_policy.py
T

319 lines
12 KiB
Python

"""Semantic Python revision-state policy helper.
Reads one JSON array of ``{"label": str, "source": str}`` records from stdin and
writes ``{"violations": [label, ...]}``. Invalid input or Python source is fatal.
"""
from __future__ import annotations
import ast
import json
import re
import string
import sys
from itertools import pairwise
from typing import Any
TARGETS = frozenset({"revision", "workspaceRevision", "selectedWorkspace"})
_FORMATTER = string.Formatter()
MAX_STATIC_TEXT = 4_096
MAX_FORMAT_SPEC = 256
MAX_STATIC_DEPTH = 64
_UNRESOLVED = object()
def _bounded_text(value: str) -> str:
if len(value) > MAX_STATIC_TEXT:
raise ValueError("static text exceeds revision policy limit")
return value
def _static_scalar(node: ast.expr, depth: int) -> object:
if depth > MAX_STATIC_DEPTH:
raise ValueError("static expression nesting exceeds revision policy limit")
if isinstance(node, ast.Constant) and type(node.value) in {
str,
int,
float,
complex,
bool,
type(None),
}:
if isinstance(node.value, str):
_bounded_text(node.value)
if isinstance(node.value, int) and node.value.bit_length() > MAX_STATIC_TEXT * 4:
raise ValueError("static integer exceeds revision policy limit")
return node.value
if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Add):
left = _static_scalar(node.left, depth + 1)
right = _static_scalar(node.right, depth + 1)
if left is _UNRESOLVED or right is _UNRESOLVED:
return _UNRESOLVED
try:
result = left + right
except TypeError:
return _UNRESOLVED
if type(result) not in {str, int, float, complex, bool}:
return _UNRESOLVED
if isinstance(result, str):
_bounded_text(result)
if isinstance(result, int) and result.bit_length() > MAX_STATIC_TEXT * 4:
raise ValueError("static integer exceeds revision policy limit")
return result
if isinstance(node, ast.JoinedStr):
result = _static_key(node, depth + 1)
return _UNRESOLVED if result is None else result
return _UNRESOLVED
def _validate_format_spec(format_spec: str) -> None:
if len(format_spec) > MAX_FORMAT_SPEC:
raise ValueError("static format specification exceeds revision policy limit")
for digits in re.findall(r"[0-9]+", format_spec):
if len(digits) > 6 or int(digits) > MAX_STATIC_TEXT:
raise ValueError("static format width or precision exceeds revision policy limit")
def _static_key(node: ast.expr, depth: int = 0) -> str | None:
if depth > MAX_STATIC_DEPTH:
raise ValueError("static key nesting exceeds revision policy limit")
if isinstance(node, ast.Constant) and isinstance(node.value, str):
return _bounded_text(node.value)
if isinstance(node, ast.BinOp) and isinstance(node.op, ast.Add):
left = _static_key(node.left, depth + 1)
right = _static_key(node.right, depth + 1)
return None if left is None or right is None else _bounded_text(left + right)
if isinstance(node, ast.JoinedStr):
pieces = []
length = 0
for value in node.values:
if isinstance(value, ast.Constant) and isinstance(value.value, str):
piece = value.value
elif isinstance(value, ast.FormattedValue):
scalar = _static_scalar(value.value, depth + 1)
if scalar is _UNRESOLVED:
return None
format_spec = "" if value.format_spec is None else _static_key(value.format_spec, depth + 1)
if format_spec is None:
return None
_validate_format_spec(format_spec)
try:
if value.conversion == ord("s"):
scalar = str(scalar)
elif value.conversion == ord("r"):
scalar = repr(scalar)
elif value.conversion == ord("a"):
scalar = ascii(scalar)
elif value.conversion != -1:
return None
piece = format(scalar, format_spec)
except (TypeError, ValueError):
return None
else:
return None
length += len(piece)
if length > MAX_STATIC_TEXT:
raise ValueError("static formatted key exceeds revision policy limit")
pieces.append(piece)
return "".join(pieces)
return None
def _is_revision_expr(node: ast.expr) -> bool:
if isinstance(node, ast.Name):
return node.id in TARGETS
if isinstance(node, ast.Attribute):
return node.attr in TARGETS
if isinstance(node, ast.Subscript):
return _static_key(node.slice) in TARGETS
return False
def _is_state_access(node: ast.AST) -> bool:
if isinstance(node, ast.Attribute):
return node.attr == "state" and _is_revision_expr(node.value)
if isinstance(node, ast.Subscript):
return _static_key(node.slice) == "state" and _is_revision_expr(node.value)
return False
def _static_sequence(node: ast.expr) -> list[ast.expr] | None:
if not isinstance(node, (ast.List, ast.Tuple)):
return None
result: list[ast.expr] = []
for element in node.elts:
if isinstance(element, ast.Starred):
nested = _static_sequence(element.value)
if nested is None:
return None
result.extend(nested)
else:
result.append(element)
return result
def _static_mapping(node: ast.expr) -> dict[str, ast.expr] | None:
if not isinstance(node, ast.Dict):
return None
result: dict[str, ast.expr] = {}
for key, value in zip(node.keys, node.values, strict=True):
if key is None:
nested = _static_mapping(value)
if nested is None:
return None
result.update(nested)
elif (name := _static_key(key)) is not None:
result[name] = value
else:
return None
return result
def _format_bindings(call: ast.Call, method: str) -> dict[str | int, ast.expr]:
if method == "format":
bindings: dict[str | int, ast.expr] = {}
position = 0
positional_known = True
for argument in call.args:
if isinstance(argument, ast.Starred):
expanded = _static_sequence(argument.value)
if expanded is None:
positional_known = False
continue
if positional_known:
for value in expanded:
bindings[position] = value
position += 1
elif positional_known:
bindings[position] = argument
position += 1
for keyword in call.keywords:
if keyword.arg is not None:
# An explicit keyword remains bound even beside **dynamic; a duplicate is TypeError.
bindings[keyword.arg] = keyword.value
else:
expanded = _static_mapping(keyword.value)
if expanded is not None:
bindings.update(expanded)
return bindings
if len(call.args) != 1 or call.keywords:
return {}
return _static_mapping(call.args[0]) or {}
def _field_accesses_state(
field_name: str, bindings: dict[str | int, ast.expr], automatic_index: int | None = None
) -> bool:
root_match = re.match(r"(?:[0-9]+|[A-Za-z_][A-Za-z0-9_]*)", field_name)
if root_match is None:
if automatic_index is None or not field_name.startswith((".", "[")):
return False
root: str | int = automatic_index
cursor = 0
else:
root_text = root_match.group(0)
root = int(root_text) if root_text.isdigit() else root_text
cursor = root_match.end()
steps: list[tuple[bool, str]] = []
while cursor < len(field_name):
if field_name[cursor] == ".":
match = re.match(r"[A-Za-z_][A-Za-z0-9_]*", field_name[cursor + 1 :])
if match is None:
return False
steps.append((True, match.group(0)))
cursor += len(match.group(0)) + 1
elif field_name[cursor] == "[":
close = field_name.find("]", cursor + 1)
if close < 0:
return False
steps.append((False, field_name[cursor + 1 : close]))
cursor = close + 1
else:
return False
if steps:
first_step = str(steps[0][1])
if str(root) in TARGETS and first_step == "state":
return True
bound = bindings.get(root)
if bound is not None and _is_revision_expr(bound) and first_step == "state":
return True
names = [str(root), *(str(key) for _is_attr, key in steps)]
return any(left in TARGETS and right == "state" for left, right in pairwise(names))
def _format_call_violation(node: ast.Call) -> bool:
function = node.func
if not isinstance(function, ast.Attribute) or function.attr not in {"format", "format_map"}:
return False
if not isinstance(function.value, ast.Constant) or not isinstance(function.value.value, str):
return False
bindings = _format_bindings(node, function.attr)
numbering: dict[str, int | str | None] = {"next": 0, "mode": None}
visited = 0
def analyze_template(template: str) -> bool:
nonlocal visited
visited += 1
if visited > 1_000:
raise ValueError("format specification nesting exceeds policy limit")
for _literal, field_name, format_spec, _conversion in _FORMATTER.parse(template):
automatic_index = None
if field_name is not None:
root_match = re.match(r"(?:[0-9]+|[A-Za-z_][A-Za-z0-9_]*)", field_name)
automatic = field_name == "" or root_match is None and field_name.startswith((".", "["))
manual = root_match is not None and root_match.group(0).isdigit()
if automatic:
if numbering["mode"] == "manual":
raise ValueError("cannot switch from manual to automatic field numbering")
numbering["mode"] = "automatic"
automatic_index = int(numbering["next"])
numbering["next"] = automatic_index + 1
elif manual:
if numbering["mode"] == "automatic":
raise ValueError("cannot switch from automatic to manual field numbering")
numbering["mode"] = "manual"
if _field_accesses_state(field_name, bindings, automatic_index):
return True
if format_spec and analyze_template(format_spec):
return True
return False
return analyze_template(function.value.value)
def has_revision_state(source: str, label: str = "<unknown>") -> bool:
tree = ast.parse(source, filename=label, mode="exec")
return any(_is_state_access(node) or (isinstance(node, ast.Call) and _format_call_violation(node)) for node in ast.walk(tree))
def analyze_batch(records: Any) -> list[str]:
if not isinstance(records, list):
raise TypeError("input must be a JSON array")
violations = []
for record in records:
if not isinstance(record, dict) or set(record) != {"label", "source"}:
raise TypeError("each record must contain exactly label and source")
label, source = record["label"], record["source"]
if not isinstance(label, str) or not isinstance(source, str):
raise TypeError("label and source must be strings")
if has_revision_state(source, label):
violations.append(label)
return violations
def main() -> int:
try:
records = json.load(sys.stdin)
json.dump({"violations": analyze_batch(records)}, sys.stdout, ensure_ascii=False)
sys.stdout.write("\n")
return 0
except Exception as error: # noqa: BLE001 - protocol boundary must fail closed
print(f"python revision-state helper failed: {error}", file=sys.stderr)
return 2
if __name__ == "__main__":
raise SystemExit(main())