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