From 19eb89b0101a1d27ddfaefa48b528f4a5e0e0538 Mon Sep 17 00:00:00 2001 From: John Kennedy <65985482+jkennedyvz@users.noreply.github.com> Date: Mon, 20 Jul 2026 13:51:46 -0700 Subject: [PATCH] Fix PR creation guard shell bypasses (#1786) Co-authored-by: langsmith-fleet[bot] (cherry picked from commit 75fb8b487852003916c4984504a13ee7226b2ceb) --- agent/middleware/pr_creation_guard.py | 61 +++++++++++++++++++++++--- tests/github/test_pr_creation_guard.py | 17 +++++++ 2 files changed, 73 insertions(+), 5 deletions(-) diff --git a/agent/middleware/pr_creation_guard.py b/agent/middleware/pr_creation_guard.py index 3420cfae..27b6429d 100644 --- a/agent/middleware/pr_creation_guard.py +++ b/agent/middleware/pr_creation_guard.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +import os import re import shlex from collections.abc import Awaitable, Callable, Mapping @@ -14,6 +15,9 @@ 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_creation_guard_shell_expansion_depth_limit__" _GITHUB_PULLS_ENDPOINT = re.compile(r"(?:^|/)repos/[^/\s]+/[^/\s]+/pulls/?$") _GITHUB_PULLS_URL = re.compile(r"https://api\.github\.com/repos/[^/\s]+/[^/\s]+/pulls/?") _BLOCK_ERROR = ( @@ -46,13 +50,59 @@ def _tool_call_id(request: ToolCallRequest) -> str | None: return None -def _shell_tokens(command: str) -> list[str]: +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:] + ): + 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 _is_assignment(token: str) -> bool: name, sep, _value = token.partition("=") return bool(sep and name and re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", name)) @@ -69,7 +119,7 @@ def _gh_subtokens(tokens: list[str], index: int) -> list[str]: def _contains_gh_pr_create(tokens: list[str]) -> bool: for index, token in enumerate(tokens): - if token != "gh": + if _executable_name(token) != "gh": continue subtokens = _gh_subtokens(tokens, index) for offset, subtoken in enumerate(subtokens[:-1]): @@ -136,7 +186,7 @@ def _gh_api_uses_post_or_body(subtokens: list[str]) -> bool: def _contains_gh_api_pull_create(tokens: list[str]) -> bool: for index, token in enumerate(tokens): - if token != "gh": + if _executable_name(token) != "gh": continue subtokens = _gh_subtokens(tokens, index) if "api" not in subtokens: @@ -153,7 +203,7 @@ def _contains_gh_api_pull_create(tokens: list[str]) -> bool: def _contains_direct_pull_create(tokens: list[str]) -> bool: for index, token in enumerate(tokens): - if token != "curl": + if _executable_name(token) != "curl": continue subtokens: list[str] = [] for candidate in tokens[index + 1 :]: @@ -193,7 +243,8 @@ def is_pr_creation_fallback_command(command: str) -> bool: """ tokens = _shell_tokens(command) return ( - _contains_gh_pr_create(tokens) + _SHELL_EXPANSION_DEPTH_LIMIT_TOKEN in tokens + or _contains_gh_pr_create(tokens) or _contains_gh_api_pull_create(tokens) or _contains_direct_pull_create(tokens) ) diff --git a/tests/github/test_pr_creation_guard.py b/tests/github/test_pr_creation_guard.py index 55b9060d..7aa7d4f0 100644 --- a/tests/github/test_pr_creation_guard.py +++ b/tests/github/test_pr_creation_guard.py @@ -38,6 +38,21 @@ def test_detects_pr_creation_fallback_commands() -> None: assert is_pr_creation_fallback_command( "curl -X POST https://api.github.com/repos/langchain-ai/open-swe/pulls -d '{}'" ) + assert is_pr_creation_fallback_command("/usr/bin/gh pr create --draft") + assert is_pr_creation_fallback_command( + "/usr/bin/curl -X POST https://api.github.com/repos/langchain-ai/open-swe/pulls -d '{}'" + ) + assert is_pr_creation_fallback_command("bash -c 'gh pr create --draft'") + assert is_pr_creation_fallback_command("GH_TOKEN=dummy sh -c 'gh pr create --draft'") + assert is_pr_creation_fallback_command( + "zsh -lc 'gh api repos/langchain-ai/open-swe/pulls -X POST -f title=x'" + ) + assert is_pr_creation_fallback_command( + "bash -c \"curl -X POST https://api.github.com/repos/langchain-ai/open-swe/pulls -d '{}'\"" + ) + assert is_pr_creation_fallback_command( + 'bash -c \'sh -c "dash -c \\"zsh -c \\\\\\"gh pr create --draft\\\\\\"\\""\'' + ) def test_allows_safe_pr_commands() -> None: @@ -45,6 +60,8 @@ def test_allows_safe_pr_commands() -> None: assert not is_pr_creation_fallback_command("gh pr list --head open-swe/foo") assert not is_pr_creation_fallback_command("gh pr edit 1 --add-label ready") assert not is_pr_creation_fallback_command("gh pr comment 1 --body done") + assert not is_pr_creation_fallback_command("bash -c 'gh pr view 1 --json url'") + assert not is_pr_creation_fallback_command("/usr/bin/gh pr view 1 --json url") async def test_middleware_blocks_execute_pr_creation_fallbacks() -> None: