"""Block shell fallbacks that submit PR review verdicts outside publish_review.""" from __future__ import annotations import json import os import re import shlex from collections.abc import Awaitable, Callable, Mapping from typing import Any from langchain.agents.middleware.types import AgentMiddleware, AgentState from langchain_core.messages import ToolMessage from langgraph.prebuilt.tool_node import ToolCallRequest from langgraph.types import Command _SHELL_SEPARATORS = {";", "&&", "||", "|", "&"} _SHELL_EXECUTABLES = {"bash", "dash", "sh", "zsh"} _MAX_SHELL_EXPANSION_DEPTH = 3 _SHELL_EXPANSION_DEPTH_LIMIT_TOKEN = "__pr_verdict_guard_shell_expansion_depth_limit__" _GITHUB_REVIEWS_ENDPOINT = re.compile( r"(?:^|/)repos/[^/\s]+/[^/\s]+/pulls/\d+/reviews(?:/\d+/events)?/?$" ) _GITHUB_REVIEWS_URL = re.compile( r"https://api\.github\.com/repos/[^/\s]+/[^/\s]+/pulls/\d+/reviews(?:/\d+/events)?/?" ) _VERDICT_EVENT_RE = re.compile(r"\bevent\b.{0,4}?(APPROVE|REQUEST_CHANGES)", re.IGNORECASE) _GH_PR_REVIEW_VERDICT_FLAGS = {"--approve", "-a", "--request-changes", "-r"} _GH_PR_REVIEW_VERDICT_PREFIXES = ("--approve=", "--request-changes=") _BLOCK_ERROR = ( "PR review verdicts (approve / request changes) must go through the " "publish_review tool, which enforces verdict authorization, self-review " "handling, and reviewer bookkeeping. Do not fall back to gh pr review, " "gh api .../reviews, curl, or another direct review-submission path. If a " "verdict was requested, call publish_review(verdict=...); otherwise " "publish a comment review." ) def _tool_name(request: ToolCallRequest) -> str | None: tool_call = getattr(request, "tool_call", None) if isinstance(tool_call, Mapping): name = tool_call.get("name") return name if isinstance(name, str) else None return None def _tool_args(request: ToolCallRequest) -> dict[str, Any]: tool_call = getattr(request, "tool_call", None) args = tool_call.get("args") if isinstance(tool_call, Mapping) else None return dict(args) if isinstance(args, Mapping) else {} def _tool_call_id(request: ToolCallRequest) -> str | None: tool_call = getattr(request, "tool_call", None) if isinstance(tool_call, Mapping): value = tool_call.get("id") return value if isinstance(value, str) else None return None def _split_shell_tokens(command: str) -> list[str]: try: return shlex.split(command, posix=True) except ValueError: return command.split() def _executable_name(token: str) -> str: return os.path.basename(token.strip("'\"")) def _shell_command_argument(tokens: list[str], shell_index: int) -> str | None: for index, token in enumerate(tokens[shell_index + 1 :], start=shell_index + 1): if token in _SHELL_SEPARATORS: return None if token == "-c" or ( token.startswith("-") and not token.startswith("--") and "c" in token[1:] ): _, _, glued = token[1:].partition("c") if glued: return glued if index + 1 < len(tokens) and tokens[index + 1] not in _SHELL_SEPARATORS: return tokens[index + 1] return None return None def _has_nested_shell_command(tokens: list[str]) -> bool: return any( _executable_name(token) in _SHELL_EXECUTABLES and _shell_command_argument(tokens, index) is not None for index, token in enumerate(tokens) ) def _expand_nested_shell_tokens(tokens: list[str], depth: int = 0) -> list[str]: if depth >= _MAX_SHELL_EXPANSION_DEPTH: if _has_nested_shell_command(tokens): return [*tokens, _SHELL_EXPANSION_DEPTH_LIMIT_TOKEN] return tokens expanded = list(tokens) for index, token in enumerate(tokens): if _executable_name(token) not in _SHELL_EXECUTABLES: continue inner_command = _shell_command_argument(tokens, index) if inner_command is None: continue expanded.extend(_expand_nested_shell_tokens(_split_shell_tokens(inner_command), depth + 1)) return expanded def _shell_tokens(command: str) -> list[str]: return _expand_nested_shell_tokens(_split_shell_tokens(command)) def _subtokens_after(tokens: list[str], index: int) -> list[str]: subtokens: list[str] = [] for token in tokens[index + 1 :]: if token in _SHELL_SEPARATORS: break subtokens.append(token) return subtokens def _contains_gh_pr_review_verdict(tokens: list[str]) -> bool: for index, token in enumerate(tokens): if _executable_name(token) != "gh": continue subtokens = _subtokens_after(tokens, index) is_pr_review = any( subtoken == "pr" and subtokens[offset + 1] == "review" for offset, subtoken in enumerate(subtokens[:-1]) ) if not is_pr_review: continue for subtoken in subtokens: if subtoken in _GH_PR_REVIEW_VERDICT_FLAGS or subtoken.startswith( _GH_PR_REVIEW_VERDICT_PREFIXES ): return True return False def _contains_gh_api_review_verdict(tokens: list[str]) -> bool: for index, token in enumerate(tokens): if _executable_name(token) != "gh": continue subtokens = _subtokens_after(tokens, index) if "api" not in subtokens: continue targets_reviews = any( _GITHUB_REVIEWS_ENDPOINT.search(subtoken.strip("'\"")) or _GITHUB_REVIEWS_URL.search(subtoken) for subtoken in subtokens ) if not targets_reviews: continue if any(_VERDICT_EVENT_RE.search(subtoken) for subtoken in subtokens): return True return False def _contains_direct_review_verdict(tokens: list[str]) -> bool: for index, token in enumerate(tokens): if _executable_name(token) != "curl": continue subtokens = _subtokens_after(tokens, index) if not any(_GITHUB_REVIEWS_URL.search(subtoken) for subtoken in subtokens): continue if any(_VERDICT_EVENT_RE.search(subtoken) for subtoken in subtokens): return True return False def is_pr_verdict_fallback_command(command: str) -> bool: """Return True when *command* submits a PR review verdict from the shell. Detection is literal-token matching — it catches the primary vectors (``gh pr review --approve/-a/--request-changes/-r``, ``gh api`` posting ``event=APPROVE|REQUEST_CHANGES`` to a ``/pulls/N/reviews`` endpoint, and ``curl`` to the reviews URL with a verdict event in the body) but is intentionally fail-open: shell aliases, ``gh`` aliases, ``--input`` JSON-file bodies, and non-curl HTTP clients will not be blocked. The ``-c`` form of a known shell (``bash``/``dash``/``sh``/``zsh``, whether space-separated or glued as ``-c'...'``) is expanded to a fixed depth and quoted/path-prefixed executables are normalized, so those do not slip through; other shells (``ash``/``busybox``/``fish``) and stdin-fed programs (``... | bash``, here-strings, ``bash -s``) remain fail-open. This is acceptable because the threat model is an honest agent papering over a ``publish_review`` limitation or refusal, not an adversary trying to bypass the guardrail. Plain ``gh pr review`` (no verdict flag) and ``--comment``/``-c`` are allowed — non-interactive ``gh`` cannot submit a verdict without one of the blocked flags. """ tokens = _shell_tokens(command) return ( _SHELL_EXPANSION_DEPTH_LIMIT_TOKEN in tokens or _contains_gh_pr_review_verdict(tokens) or _contains_gh_api_review_verdict(tokens) or _contains_direct_review_verdict(tokens) ) def _blocked_tool_message(request: ToolCallRequest, command: str) -> ToolMessage: content = { "status": "error", "error_type": "PullRequestVerdictFallbackBlocked", "code": "pr_verdict_fallback_blocked", "recoverable_by_agent": False, "error": _BLOCK_ERROR, "blocked_command": command, } return ToolMessage( content=json.dumps(content), tool_call_id=_tool_call_id(request), status="error", ) class PullRequestVerdictGuardMiddleware(AgentMiddleware): """Keep APPROVE/REQUEST_CHANGES submissions on the guarded publish_review path.""" state_schema = AgentState def _blocked_message_for_request(self, request: ToolCallRequest) -> ToolMessage | None: if _tool_name(request) != "execute": return None command = _tool_args(request).get("command") if not isinstance(command, str) or not is_pr_verdict_fallback_command(command): return None return _blocked_tool_message(request, command) def wrap_tool_call( self, request: ToolCallRequest, handler: Callable[[ToolCallRequest], ToolMessage | Command], ) -> ToolMessage | Command: blocked = self._blocked_message_for_request(request) if blocked is not None: return blocked return handler(request) async def awrap_tool_call( self, request: ToolCallRequest, handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]], ) -> ToolMessage | Command: blocked = self._blocked_message_for_request(request) if blocked is not None: return blocked return await handler(request)