diff --git a/agent/middleware/__init__.py b/agent/middleware/__init__.py index 19bc0466..8b87afae 100644 --- a/agent/middleware/__init__.py +++ b/agent/middleware/__init__.py @@ -9,6 +9,7 @@ _MIDDLEWARE_MODULES = { "ModelFallbackMiddleware": ".model_fallback", "notify_step_limit_reached": ".notify_step_limit", "PlanModeMiddleware": ".plan_mode", + "PullRequestCreationGuardMiddleware": ".pr_creation_guard", "refresh_github_proxy_before_model": ".refresh_github_proxy", "RepairOrphanedToolCallsMiddleware": ".repair_orphaned_tool_calls", "SlackAssistantStatusMiddleware": ".refresh_slack_status", @@ -31,6 +32,7 @@ __all__ = [ "ExcludeToolsMiddleware", "ModelFallbackMiddleware", "PlanModeMiddleware", + "PullRequestCreationGuardMiddleware", "RepairOrphanedToolCallsMiddleware", "SanitizeFireworksMessagesMiddleware", "SanitizeOpenAIResponsesMiddleware", @@ -59,6 +61,7 @@ if TYPE_CHECKING: from .model_fallback import ModelFallbackMiddleware from .notify_step_limit import notify_step_limit_reached from .plan_mode import PlanModeMiddleware + from .pr_creation_guard import PullRequestCreationGuardMiddleware from .refresh_github_proxy import refresh_github_proxy_before_model from .refresh_slack_status import SlackAssistantStatusMiddleware from .repair_orphaned_tool_calls import RepairOrphanedToolCallsMiddleware diff --git a/agent/middleware/pr_creation_guard.py b/agent/middleware/pr_creation_guard.py new file mode 100644 index 00000000..3420cfae --- /dev/null +++ b/agent/middleware/pr_creation_guard.py @@ -0,0 +1,249 @@ +"""Block shell fallbacks that create pull requests outside open_pull_request.""" + +from __future__ import annotations + +import json +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 = {";", "&&", "||", "|", "&"} +_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 = ( + "New pull requests must be opened with the open_pull_request tool so the PR is " + "attributed to the triggering user. If open_pull_request failed, surface that " + "failure instead of falling back to gh pr create, gh api /pulls, curl, or another " + "direct PR creation path." +) + + +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 _shell_tokens(command: str) -> list[str]: + try: + return shlex.split(command, posix=True) + except ValueError: + return command.split() + + +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)) + + +def _gh_subtokens(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_create(tokens: list[str]) -> bool: + for index, token in enumerate(tokens): + if token != "gh": + continue + subtokens = _gh_subtokens(tokens, index) + for offset, subtoken in enumerate(subtokens[:-1]): + if subtoken == "pr" and subtokens[offset + 1] == "create": + return True + return False + + +_GH_API_VALUE_FLAGS = { + "-X", + "--method", + "-H", + "--header", + "-F", + "--field", + "-f", + "--raw-field", + "--hostname", + "--input", + "-q", + "--jq", + "-p", + "--preview", + "--cache", + "-t", + "--template", +} + + +def _gh_api_endpoint(subtokens: list[str]) -> str | None: + for index, token in enumerate(subtokens): + if token != "api": + continue + skip_next = False + for candidate in subtokens[index + 1 :]: + if skip_next: + skip_next = False + continue + if candidate.startswith("-"): + if "=" not in candidate and candidate in _GH_API_VALUE_FLAGS: + skip_next = True + continue + if _is_assignment(candidate): + continue + return candidate.strip("'\"") + return None + + +def _gh_api_uses_post_or_body(subtokens: list[str]) -> bool: + body_flags = {"-f", "--field", "-F", "--raw-field", "--input"} + for index, token in enumerate(subtokens): + upper = token.upper() + if upper in {"-XPOST", "--METHOD=POST"}: + return True + if token in {"-X", "--method"} and index + 1 < len(subtokens): + if subtokens[index + 1].upper() == "POST": + return True + if token.startswith("--method=") and token.split("=", 1)[1].upper() == "POST": + return True + if token in body_flags or any(token.startswith(f"{flag}=") for flag in body_flags): + return True + return False + + +def _contains_gh_api_pull_create(tokens: list[str]) -> bool: + for index, token in enumerate(tokens): + if token != "gh": + continue + subtokens = _gh_subtokens(tokens, index) + if "api" not in subtokens: + continue + endpoint = _gh_api_endpoint(subtokens) + if ( + endpoint + and _GITHUB_PULLS_ENDPOINT.search(endpoint) + and _gh_api_uses_post_or_body(subtokens) + ): + return True + return False + + +def _contains_direct_pull_create(tokens: list[str]) -> bool: + for index, token in enumerate(tokens): + if token != "curl": + continue + subtokens: list[str] = [] + for candidate in tokens[index + 1 :]: + if candidate in _SHELL_SEPARATORS: + break + subtokens.append(candidate) + if not any(_GITHUB_PULLS_URL.search(candidate) for candidate in subtokens): + continue + has_post = any( + token.upper() in {"-XPOST", "--REQUEST=POST"} + or ( + token in {"-X", "--request"} + and idx + 1 < len(subtokens) + and subtokens[idx + 1].upper() == "POST" + ) + or (token.startswith("--request=") and token.split("=", 1)[1].upper() == "POST") + for idx, token in enumerate(subtokens) + ) + has_body = any(token in {"-d", "--data", "--data-raw", "--json"} for token in subtokens) + if has_post or has_body: + return True + return False + + +def is_pr_creation_fallback_command(command: str) -> bool: + """Return True when *command* is a known PR-creation shell fallback. + + Detection is literal-token matching — it catches the primary vectors + (``gh pr create``, ``gh api repos/.../pulls -X POST``, ``curl`` to + ``/pulls``) but is intentionally fail-open: shell aliases, ``gh`` + aliases (``gh prc`` set via ``gh alias set``), flags between ``pr`` and + ``create`` (``gh pr --repo x create``), and non-curl HTTP clients + (``python -c …``, ``wget``, etc.) will not be blocked. This is + acceptable because the threat model is an honest agent papering over an + ``open_pull_request`` failure, not an adversary trying to bypass the + guardrail. + """ + tokens = _shell_tokens(command) + return ( + _contains_gh_pr_create(tokens) + or _contains_gh_api_pull_create(tokens) + or _contains_direct_pull_create(tokens) + ) + + +def _blocked_tool_message(request: ToolCallRequest, command: str) -> ToolMessage: + content = { + "status": "error", + "error_type": "PullRequestCreationFallbackBlocked", + "code": "pr_creation_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 PullRequestCreationGuardMiddleware(AgentMiddleware): + """Prevent attributed-PR failures from being hidden by shell fallbacks.""" + + 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_creation_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) diff --git a/agent/prompt.py b/agent/prompt.py index 450acce8..6bc2c737 100644 --- a/agent/prompt.py +++ b/agent/prompt.py @@ -307,7 +307,7 @@ Steps, in order: **IMPORTANT: Never force-push.** Never run `git push --force` or `git push --force-with-lease`, and never amend or rebase commits that are already on the remote branch — reviewers rely on inter-commit diffs. Add follow-up work as new commits. If a normal push is rejected because the remote branch has new commits, run `git pull --rebase origin ` and push again; if that conflicts, report it and stop. -**IMPORTANT: If `git push`, `open_pull_request`, or `gh pr edit` fails with an infrastructure or permission error, do not retry blindly. Report the failure and end the task.** +**IMPORTANT: If `git push`, `open_pull_request`, or `gh pr edit` fails with an infrastructure/permission/access error — including "403", "404"/"Not Found" from `open_pull_request`, "GitHub App not installed/access denied", or "Permission denied" — do not retry via `gh pr create`, `gh api repos/.../pulls`, direct REST `POST /repos/.../pulls`, or any other PR creation fallback. Report the failure to the user and end the task.** **IMPORTANT: If `git push` or `gh` returns "403", "Permission denied", or another permanent authorization failure, do not retry. Report the error to the user immediately and stop.** diff --git a/agent/server.py b/agent/server.py index 72658ec7..aec5ecb5 100644 --- a/agent/server.py +++ b/agent/server.py @@ -64,6 +64,7 @@ from .integrations.notion_mcp import load_notion_tools from .middleware import ( ModelFallbackMiddleware, PlanModeMiddleware, + PullRequestCreationGuardMiddleware, SandboxCircuitBreakerMiddleware, SanitizeFireworksMessagesMiddleware, SanitizeOpenAIResponsesMiddleware, @@ -1031,6 +1032,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: max_delay=10.0, ), ToolArtifactMiddleware(), + PullRequestCreationGuardMiddleware(), WorkflowPushGuardMiddleware(), refresh_github_proxy_before_model, check_message_queue_before_model, diff --git a/agent/tools/open_pull_request.py b/agent/tools/open_pull_request.py index 2abfff3e..1389bcb3 100644 --- a/agent/tools/open_pull_request.py +++ b/agent/tools/open_pull_request.py @@ -4,6 +4,7 @@ from __future__ import annotations import logging from typing import Any +from urllib.parse import quote import httpx from langgraph.config import get_config @@ -21,6 +22,9 @@ logger = logging.getLogger(__name__) GITHUB_API = "https://api.github.com" _USER_TOKEN_SOURCES = ("slack", "dashboard") _REFERENCES_HEADING = "## References" +_ACCESS_FAILURE_CODE = "github_app_access_missing_or_repo_not_found" +_BRANCH_FAILURE_CODE = "github_pr_branch_not_visible" +_PREFLIGHT_FAILURE_CODE = "github_pr_preflight_failed" async def _resolve_pr_author_token() -> tuple[str | None, str]: @@ -67,6 +71,360 @@ def _auth_headers(token: str) -> dict[str, str]: } +def _github_message(resp: httpx.Response) -> str: + try: + data = resp.json() + except Exception: + return resp.text.strip() or f"HTTP {resp.status_code}" + if isinstance(data, dict): + message = data.get("message") + if isinstance(message, str) and message.strip(): + return message.strip() + return resp.text.strip() or f"HTTP {resp.status_code}" + + +def _configurable() -> dict[str, Any]: + try: + config = get_config() + except Exception: + return {} + configurable = config.get("configurable", {}) if isinstance(config, dict) else {} + return dict(configurable) if isinstance(configurable, dict) else {} + + +def _head_branch_for_repo(owner: str, head: str) -> str | None: + if ":" not in head: + return head + head_owner, branch = head.split(":", 1) + if head_owner == owner and branch: + return branch + return None + + +def _failure_payload( + *, + code: str, + owner: str, + repo: str, + head: str, + base: str, + token_kind: str, + http_status: int | None, + reason: str, + likely_cause: str, + suggested_action: str, + branch_pushed: bool | None, + failed_step: str, + repo_visible: bool | None = None, + base_branch_visible: bool | None = None, + head_branch_visible: bool | None = None, +) -> dict[str, Any]: + error = ( + "Failed to open an attributed PR with open_pull_request. " + f"Reason: {reason}. Likely cause: {likely_cause}. " + f"Branch pushed: {owner}/{repo}:{head} " + f"({'unknown' if branch_pushed is None else 'yes' if branch_pushed else 'no'}). " + "PR created: no. " + f"Action: {suggested_action}" + ) + payload: dict[str, Any] = { + "success": False, + "error": error, + "code": code, + "recoverable_by_agent": False, + "owner": owner, + "repo": repo, + "head": head, + "base": base, + "token_kind": token_kind, + "http_status": http_status, + "branch_pushed": branch_pushed, + "pr_created": False, + "failed_step": failed_step, + "likely_cause": likely_cause, + "suggested_action": suggested_action, + } + if repo_visible is not None: + payload["repo_visible"] = repo_visible + if base_branch_visible is not None: + payload["base_branch_visible"] = base_branch_visible + if head_branch_visible is not None: + payload["head_branch_visible"] = head_branch_visible + _record_open_pr_failure_telemetry(payload) + return payload + + +def _record_open_pr_failure_telemetry(payload: dict[str, Any]) -> None: + configurable = _configurable() + logger.warning( + "open_pull_request_failed code=%s owner=%s repo=%s head=%s base=%s " + "http_status=%s token_kind=%s branch_pushed=%s thread_id=%s source=%s", + payload.get("code"), + payload.get("owner"), + payload.get("repo"), + payload.get("head"), + payload.get("base"), + payload.get("http_status"), + payload.get("token_kind"), + payload.get("branch_pushed"), + configurable.get("thread_id"), + configurable.get("source"), + extra={ + "open_pull_request_failure": { + "code": payload.get("code"), + "owner": payload.get("owner"), + "repo": payload.get("repo"), + "head": payload.get("head"), + "base": payload.get("base"), + "http_status": payload.get("http_status"), + "token_kind": payload.get("token_kind"), + "branch_pushed": payload.get("branch_pushed"), + "pr_created": payload.get("pr_created"), + "failed_step": payload.get("failed_step"), + "thread_id": configurable.get("thread_id"), + "source": configurable.get("source"), + } + }, + ) + + +def _access_failure_payload( + *, + owner: str, + repo: str, + head: str, + base: str, + token_kind: str, + http_status: int | None, + reason: str, + branch_pushed: bool | None, + failed_step: str, + repo_visible: bool | None = None, + base_branch_visible: bool | None = None, + head_branch_visible: bool | None = None, +) -> dict[str, Any]: + return _failure_payload( + code=_ACCESS_FAILURE_CODE, + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=http_status, + reason=reason, + likely_cause=( + "the Open SWE GitHub App or PR author token is not installed on, granted access " + "to, or able to see this repository or one of the PR branches" + ), + suggested_action=( + "install or grant the Open SWE GitHub App and the triggering user's GitHub " + "authorization access to this repository, verify the base/head branches exist, " + "then ask Open SWE to retry opening the PR" + ), + branch_pushed=branch_pushed, + failed_step=failed_step, + repo_visible=repo_visible, + base_branch_visible=base_branch_visible, + head_branch_visible=head_branch_visible, + ) + + +def _branch_failure_payload( + *, + owner: str, + repo: str, + head: str, + base: str, + token_kind: str, + http_status: int, + branch: str, + branch_role: str, +) -> dict[str, Any]: + branch_pushed = False if branch_role == "head" else None + return _failure_payload( + code=_BRANCH_FAILURE_CODE, + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=http_status, + reason=f"GitHub could not see the {branch_role} branch `{branch}` before PR creation", + likely_cause=( + f"the {branch_role} branch does not exist on `{owner}/{repo}` or is not visible " + "to the PR author token" + ), + suggested_action=( + f"push or restore the {branch_role} branch `{branch}`, ensure the Open SWE " + "GitHub App/token can see it, then ask Open SWE to retry opening the PR" + ), + branch_pushed=branch_pushed, + failed_step=f"preflight_{branch_role}_branch", + repo_visible=True, + base_branch_visible=False if branch_role == "base" else True, + head_branch_visible=False if branch_role == "head" else None, + ) + + +async def _github_get(client: httpx.AsyncClient, token: str, path: str) -> httpx.Response: + return await client.get(f"{GITHUB_API}{path}", headers=_auth_headers(token)) + + +async def _preflight_pr_access( + *, + client: httpx.AsyncClient, + token: str, + token_kind: str, + owner: str, + repo: str, + head: str, + base: str, +) -> dict[str, Any] | None: + """Diagnose *why* a PR creation POST failed — not an authoritative gate. + + This runs only after the POST returns a non-201/non-422 status so it + doesn't add latency on the happy path and avoids false positives from + read-after-write inconsistency on just-pushed head branches. + """ + repo_resp = await _github_get(client, token, f"/repos/{owner}/{repo}") + if repo_resp.status_code in {403, 404}: + return _access_failure_payload( + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=repo_resp.status_code, + reason=f"GitHub returned {repo_resp.status_code} while checking repository access", + branch_pushed=None, + failed_step="preflight_repo", + repo_visible=False, + ) + if repo_resp.status_code != 200: + return _failure_payload( + code=_PREFLIGHT_FAILURE_CODE, + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=repo_resp.status_code, + reason=( + f"GitHub returned {repo_resp.status_code} while checking repository access: " + f"{_github_message(repo_resp)}" + ), + likely_cause="GitHub repository access preflight failed before PR creation", + suggested_action="check GitHub availability and repository access, then retry", + branch_pushed=None, + failed_step="preflight_repo", + repo_visible=None, + ) + + base_resp = await _github_get( + client, token, f"/repos/{owner}/{repo}/branches/{quote(base, safe='')}" + ) + if base_resp.status_code == 404: + return _branch_failure_payload( + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=base_resp.status_code, + branch=base, + branch_role="base", + ) + if base_resp.status_code in {401, 403}: + return _access_failure_payload( + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=base_resp.status_code, + reason=f"GitHub returned {base_resp.status_code} while checking base branch access", + branch_pushed=None, + failed_step="preflight_base_branch", + repo_visible=True, + base_branch_visible=False, + ) + if base_resp.status_code != 200: + return _failure_payload( + code=_PREFLIGHT_FAILURE_CODE, + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=base_resp.status_code, + reason=( + f"GitHub returned {base_resp.status_code} while checking base branch access: " + f"{_github_message(base_resp)}" + ), + likely_cause="GitHub branch access preflight failed before PR creation", + suggested_action="check GitHub availability and branch access, then retry", + branch_pushed=None, + failed_step="preflight_base_branch", + repo_visible=True, + base_branch_visible=None, + ) + + head_branch = _head_branch_for_repo(owner, head) + if head_branch is None: + return None + head_resp = await _github_get( + client, token, f"/repos/{owner}/{repo}/branches/{quote(head_branch, safe='')}" + ) + if head_resp.status_code == 404: + return _branch_failure_payload( + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=head_resp.status_code, + branch=head_branch, + branch_role="head", + ) + if head_resp.status_code in {401, 403}: + return _access_failure_payload( + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=head_resp.status_code, + reason=f"GitHub returned {head_resp.status_code} while checking head branch access", + branch_pushed=False, + failed_step="preflight_head_branch", + repo_visible=True, + base_branch_visible=True, + head_branch_visible=False, + ) + if head_resp.status_code != 200: + return _failure_payload( + code=_PREFLIGHT_FAILURE_CODE, + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=token_kind, + http_status=head_resp.status_code, + reason=( + f"GitHub returned {head_resp.status_code} while checking head branch access: " + f"{_github_message(head_resp)}" + ), + likely_cause="GitHub branch access preflight failed before PR creation", + suggested_action="check GitHub availability and branch access, then retry", + branch_pushed=None, + failed_step="preflight_head_branch", + repo_visible=True, + base_branch_visible=True, + head_branch_visible=None, + ) + return None + + async def _find_existing_pr( client: httpx.AsyncClient, token: str, owner: str, repo: str, head: str ) -> dict[str, Any] | None: @@ -268,10 +626,20 @@ async def _open_pull_request( ) -> dict[str, Any]: token, kind = await _resolve_pr_author_token() if not token: - return { - "success": False, - "error": "No GitHub token available to open the pull request.", - } + return _failure_payload( + code="no_github_token", + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=kind, + http_status=None, + reason="No GitHub token was available to open the pull request", + likely_cause="the triggering user is not authorized and no GitHub App token is available", + suggested_action="connect GitHub authorization or install/grant the Open SWE GitHub App, then retry", + branch_pushed=None, + failed_step="resolve_pr_author_token", + ) async with httpx.AsyncClient(timeout=30.0) as client: body = await _maybe_append_references(client, token, owner, repo, body) @@ -325,10 +693,60 @@ async def _open_pull_request( "token_kind": kind, } - return { - "success": False, - "error": f"GitHub returned {resp.status_code}: {resp.text}", - } + # POST failed — run preflight diagnostics to understand why, so the + # agent gets an actionable diagnosis instead of a raw HTTP status. + # Preflight runs *after* the POST so a just-pushed branch that is + # momentarily invisible to GitHub's ref endpoints does not cause a + # false-positive preflight failure on what would have been a successful + # PR creation. + diagnostic = await _preflight_pr_access( + client=client, + token=token, + token_kind=kind, + owner=owner, + repo=repo, + head=head, + base=base, + ) + if diagnostic is not None: + return diagnostic + + # Preflight found nothing — surface the raw POST failure. + if resp.status_code == 404: + return _access_failure_payload( + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=kind, + http_status=resp.status_code, + reason="GitHub returned 404 while creating the pull request", + branch_pushed=True, + failed_step="create_pull_request", + repo_visible=True, + base_branch_visible=True, + head_branch_visible=True + if _head_branch_for_repo(owner, head) is not None + else None, + ) + + return _failure_payload( + code="github_pr_create_failed", + owner=owner, + repo=repo, + head=head, + base=base, + token_kind=kind, + http_status=resp.status_code, + reason=f"GitHub returned {resp.status_code} while creating the pull request: {_github_message(resp)}", + likely_cause="GitHub rejected the pull request creation request", + suggested_action="inspect the GitHub error, correct the branch or repository state, then retry", + branch_pushed=True, + failed_step="create_pull_request", + repo_visible=True, + base_branch_visible=True, + head_branch_visible=True if _head_branch_for_repo(owner, head) is not None else None, + ) async def open_pull_request( @@ -363,7 +781,8 @@ async def open_pull_request( Returns: On success: {"success": True, "created": bool, "url": str, "number": int, "author": str}. ``created`` is False when an open PR already existed. - On failure: {"success": False, "error": str}. + On failure: {"success": False, "error": str, "code": str, + "recoverable_by_agent": False, "pr_created": False, ...}. """ return await _open_pull_request( owner=owner, diff --git a/docs/upstream-sync/triage.jsonl b/docs/upstream-sync/triage.jsonl index cb7d05b4..dbfe1cc3 100644 --- a/docs/upstream-sync/triage.jsonl +++ b/docs/upstream-sync/triage.jsonl @@ -68,7 +68,7 @@ {"sha": "52fe2916", "pr": 1698, "subject": "feat: add PR review link route (#1698)", "disposition": "landed", "reason": "cherry-picked (-x) in #127 (PR review link route)", "branch": "reviewer-misc", "local_sha": null, "updated": "2026-07-08T22:58:21Z"} {"sha": "5f7f5fbd", "pr": 1697, "subject": "fix: Reduce graph import and loader startup latency (#1697)", "disposition": "landed", "reason": "import-hygiene refactor; cross-cutting, references many deferred upstream-only modules", "branch": "durable-dispatch", "local_sha": null, "updated": "2026-07-09T17:22:53Z"} {"sha": "216cf181", "pr": 1699, "subject": "fix: keep workflow HITL without token downscoping (#1699)", "disposition": "landed", "reason": "DIVERGES-FROM-UPSTREAM: fork deliberately does NOT adopt #1699's standing-token workflows:write broadening. Security review (#159) BLOCKed it — the standing ALWAYS-ON proxy token carrying workflows:write turns the HITL guard's git-push-parser gaps (obfuscated-expansion push, `gh api` REST contents PUT, cross-branch refspecs) into live unapproved-workflow-push exploits. Fork keeps BASE without workflows:write and restores the transient per-approval elevation (_run_with_workflow_token mints WORKFLOW_RUNTIME_PROXY_TOKEN_PERMISSIONS around the approved, guard-normalized fixed_command, then downscopes to RUNTIME then BASE): the token scope is the backstop the parser relies on, so a bypass hits GitHub 403. HITL diff-preview/approval-URL/Slack-card additions from #159 retained; token-model divergence only.", "branch": "plan-approval", "local_sha": null, "updated": "2026-07-09T20:00:00Z"} -{"sha": "67abf5b0", "pr": 1659, "subject": "fix: surface attributed PR creation failures (#1659)", "disposition": "deferred", "reason": "PR-attribution-failure guard (new mw, safe imports); heavy conflict on diverged open_pull_request.py", "branch": "pr-attribution", "local_sha": null, "updated": "2026-07-08T20:14:42Z"} +{"sha": "67abf5b0", "pr": 1659, "subject": "fix: surface attributed PR creation failures (#1659)", "disposition": "landed", "reason": "PR-attribution-failure guard (new mw, safe imports); heavy conflict on diverged open_pull_request.py", "branch": "pr-attribution", "local_sha": null, "updated": "2026-07-13T17:16:20Z"} {"sha": "3dbc0282", "pr": 1676, "subject": "fix: preserve plan redirects after login (#1676)", "disposition": "landed", "reason": "FLAG-HUMAN: follow-on to landed #1668 refining sanitize_redirect_to (open-redirect auth surface); not a dup", "branch": "plan-approval", "local_sha": null, "updated": "2026-07-09T17:10:22Z"} {"sha": "c75cbb1f", "pr": 1677, "subject": "feat: re-add Fable 5 with an admin toggle to disable it (#1677)", "disposition": "landed", "reason": "Ported + Bedrock-converted onto dev via feat/readd-fable5-bedrock; anthropic: Fable ID mapped to bedrock_converse:us.anthropic.claude-fable-5.", "branch": "fable-admin-toggle", "local_sha": null, "updated": "2026-07-10T17:07:19Z"} {"sha": "bb104d93", "pr": 1679, "subject": "fix: submit plan comments with cmd enter (#1679)", "disposition": "landed", "reason": "applies clean but edits fork-diverged PlanReview.tsx (#130); needs UI/e2e validation — separate PR", "branch": "plan-approval", "local_sha": null, "updated": "2026-07-09T17:10:22Z"} diff --git a/docs/upstream-sync/triage.md b/docs/upstream-sync/triage.md index 83cef60f..9bb0ea1c 100644 --- a/docs/upstream-sync/triage.md +++ b/docs/upstream-sync/triage.md @@ -53,6 +53,7 @@ Rows key on the **upstream SHA** (stable across local cherry-picks). Deferred ro | `52fe2916` | #1698 | feat: add PR review link route (#1698) | Landed | cherry-picked (-x) in #127 (PR review link route) | reviewer-misc | | `5f7f5fbd` | #1697 | fix: Reduce graph import and loader startup latency (#1697) | Landed | import-hygiene refactor; cross-cutting, references many deferred upstream-only modules | durable-dispatch | | `216cf181` | #1699 | fix: keep workflow HITL without token downscoping (#1699) | Landed | DIVERGES-FROM-UPSTREAM: fork deliberately does NOT adopt #1699's standing-token workflows:write broadening. Security review (#159) BLOCKed it — the standing ALWAYS-ON proxy token carrying workflows:write turns the HITL guard's git-push-parser gaps (obfuscated-expansion push, `gh api` REST contents PUT, cross-branch refspecs) into live unapproved-workflow-push exploits. Fork keeps BASE without workflows:write and restores the transient per-approval elevation (_run_with_workflow_token mints WORKFLOW_RUNTIME_PROXY_TOKEN_PERMISSIONS around the approved, guard-normalized fixed_command, then downscopes to RUNTIME then BASE): the token scope is the backstop the parser relies on, so a bypass hits GitHub 403. HITL diff-preview/approval-URL/Slack-card additions from #159 retained; token-model divergence only. | plan-approval | +| `67abf5b0` | #1659 | fix: surface attributed PR creation failures (#1659) | Landed | PR-attribution-failure guard (new mw, safe imports); heavy conflict on diverged open_pull_request.py | pr-attribution | | `3dbc0282` | #1676 | fix: preserve plan redirects after login (#1676) | Landed | FLAG-HUMAN: follow-on to landed #1668 refining sanitize_redirect_to (open-redirect auth surface); not a dup | plan-approval | | `c75cbb1f` | #1677 | feat: re-add Fable 5 with an admin toggle to disable it (#1677) | Landed | Ported + Bedrock-converted onto dev via feat/readd-fable5-bedrock; anthropic: Fable ID mapped to bedrock_converse:us.anthropic.claude-fable-5. | fable-admin-toggle | | `bb104d93` | #1679 | fix: submit plan comments with cmd enter (#1679) | Landed | applies clean but edits fork-diverged PlanReview.tsx (#130); needs UI/e2e validation — separate PR | plan-approval | @@ -91,7 +92,6 @@ Rows key on the **upstream SHA** (stable across local cherry-picks). Deferred ro | `c0a7e93e` | #1691 | fix: reconnect sandbox backend on resumed runs (#1691) | Deferred | reconnect proxy (has_backend/reconnect); assumes async create_sandbox | sandbox-refactor | | `4f8bc2dd` | #1692 | refactor: simplify open-swe agent sandbox lifecycle (#1692) | Deferred | FLAG-HUMAN: structural rewrite of ensure_sandbox_for_thread (drops __creating__ 4-case sentinel) | sandbox-refactor | | `48217b68` | #1489 | feat(open-swe): add E2B sandbox provider (#1489) | Deferred | additive E2B provider; separable but ships on the async sandbox.py base | sandbox-refactor | -| `67abf5b0` | #1659 | fix: surface attributed PR creation failures (#1659) | Deferred | PR-attribution-failure guard (new mw, safe imports); heavy conflict on diverged open_pull_request.py | pr-attribution | | `5003c953` | #1683 | feat: open Linear-triggered PRs as the triggering user (#1683) | Deferred | FLAG-HUMAN: adds linear to resolve_github_token per-user OAuth branch (auth surface); depends on #1626 linear.py | linear-pr-as-user | | `22e024cb` | #1704 | fix: link issue PRs and prompt repo conventions (#1704) | Deferred | issue/PR linking + repo-convention prompt; clean but prompt-conflict risk vs #113 | webhook-issue-linking | | `27b0ddeb` | #1708 | feat: add GPT-5.6 OpenAI models (#1708) | Deferred | FLAG-HUMAN: adds OpenAI GPT-5.6 to the model picker; fork's picker is Bedrock/Fireworks-only — needs a product decision before adopting OpenAI models. Gateway (#155) can route OpenAI if adopted. | model-picker | diff --git a/tests/e2e/fakes.py b/tests/e2e/fakes.py index b1dde372..ba57ee82 100644 --- a/tests/e2e/fakes.py +++ b/tests/e2e/fakes.py @@ -124,6 +124,15 @@ def _diff_files(base: str, head: str) -> list[dict[str, Any]]: return files +def branch_exists(branch: str) -> bool: + """Check whether a branch exists in the bare remote (the fake GitHub).""" + try: + _git("--git-dir", str(BARE_REMOTE), "rev-parse", "--verify", f"refs/heads/{branch}") + return True + except subprocess.CalledProcessError: + return False + + def create_pull( owner: str, repo: str, *, head: str, base: str, title: str, body: str, draft: bool ) -> dict[str, Any]: diff --git a/tests/e2e/harness.py b/tests/e2e/harness.py index f774d546..cc7c9dec 100644 --- a/tests/e2e/harness.py +++ b/tests/e2e/harness.py @@ -427,6 +427,13 @@ async def gh_get_repo(owner: str, repo: str) -> JSONResponse: return JSONResponse({"full_name": f"{owner}/{repo}", "private": False}) +@app.get("/fake-gh/repos/{owner}/{repo}/branches/{branch:path}") +async def gh_get_branch(owner: str, repo: str, branch: str) -> JSONResponse: # noqa: ARG001 + if not fakes.branch_exists(branch): + return JSONResponse({"message": "Branch not found"}, status_code=404) + return JSONResponse({"name": branch, "commit": {"sha": "deadbeef"}}) + + @app.get("/fake-gh/repos/{owner}/{repo}/pulls") async def gh_list_pulls(owner: str, repo: str) -> JSONResponse: # noqa: ARG001 return JSONResponse([]) diff --git a/tests/test_github_comment_prompts.py b/tests/test_github_comment_prompts.py index e08e4e01..df783611 100644 --- a/tests/test_github_comment_prompts.py +++ b/tests/test_github_comment_prompts.py @@ -185,6 +185,15 @@ def test_construct_system_prompt_forbids_force_push() -> None: assert "git pull --rebase origin " in prompt +def test_construct_system_prompt_forbids_pr_creation_fallbacks() -> None: + prompt = construct_system_prompt(working_dir="/workspace") + + assert '"404"/"Not Found" from `open_pull_request`' in prompt + assert "do not retry via `gh pr create`" in prompt + assert "`gh api repos/.../pulls`" in prompt + assert "direct REST `POST /repos/.../pulls`" in prompt + + def test_construct_system_prompt_emits_no_attribution_when_identity_present() -> None: identity = CollaboratorIdentity( display_name="octocat", diff --git a/tests/test_open_pull_request.py b/tests/test_open_pull_request.py index a8a5c0f1..e30a88fd 100644 --- a/tests/test_open_pull_request.py +++ b/tests/test_open_pull_request.py @@ -59,11 +59,12 @@ class _FakeClient: return self._post async def get( - self, url: str, *, headers: dict[str, str], params: dict[str, str] + self, url: str, *, headers: dict[str, str], params: dict[str, str] | None = None ) -> _FakeResponse: self.get_calls.append({"url": url, "headers": headers, "params": params}) - assert self._get is not None - return self._get + if self._get is not None: + return self._get + return _FakeResponse(200, {"name": "ok"}) class _RoutingClient: @@ -94,7 +95,7 @@ class _RoutingClient: for needle, resp in self._get_routes.items(): if needle in url: return resp - raise AssertionError(f"unexpected GET {url}") + return _FakeResponse(200, {"name": "ok"}) def _install_client(monkeypatch: pytest.MonkeyPatch, client: _FakeClient | _RoutingClient) -> None: @@ -265,7 +266,8 @@ def test_returns_existing_pr_on_422(monkeypatch: pytest.MonkeyPatch) -> None: assert result["success"] is True assert result["created"] is False assert result["number"] == 9 - assert client.get_calls[0]["params"] == { + pr_lookup = [call for call in client.get_calls if call["params"]] + assert pr_lookup[0]["params"] == { "head": "langchain-ai:open-swe/feature", "state": "open", } @@ -279,15 +281,74 @@ def test_error_surfaced_on_failure(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(profiles, "get_valid_access_token", lambda *_a, **_k: _coro("user-tok")) monkeypatch.setattr(opr, "get_github_app_installation_token", lambda: _coro("bot")) - client = _FakeClient(post=_FakeResponse(403, text="Resource not accessible")) + client = _FakeClient(post=_FakeResponse(403, {"message": "Resource not accessible"})) _install_client(monkeypatch, client) result = _open() assert result["success"] is False + assert result["code"] == "github_pr_create_failed" + assert result["recoverable_by_agent"] is False + assert result["pr_created"] is False assert "403" in result["error"] +def test_404_create_returns_actionable_access_diagnostic( + monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + _set_config(monkeypatch, {"source": "slack", "github_login": "johannes117", "thread_id": "t1"}) + _stub_token(monkeypatch) + client = _FakeClient(post=_FakeResponse(404, {"message": "Not Found"})) + _install_client(monkeypatch, client) + + result = _open() + + assert result["success"] is False + assert result["code"] == "github_app_access_missing_or_repo_not_found" + assert result["recoverable_by_agent"] is False + assert result["owner"] == "langchain-ai" + assert result["repo"] == "open-swe" + assert result["head"] == "open-swe/feature" + assert result["base"] == "main" + assert result["branch_pushed"] is True + assert result["pr_created"] is False + assert "install or grant" in result["suggested_action"] + assert "PR created: no" in result["error"] + assert ( + "open_pull_request_failed code=github_app_access_missing_or_repo_not_found" in caplog.text + ) + + +def test_preflight_head_branch_404_reports_branch_not_pushed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Preflight runs as a diagnostic after the POST fails — the POST is attempted + first so a just-pushed branch that momentarily 404s from GitHub's ref + endpoints doesn't cause a false-positive preflight failure.""" + _set_config(monkeypatch, {"source": "slack", "github_login": "johannes117"}) + _stub_token(monkeypatch) + client = _RoutingClient( + post=_FakeResponse(422, {"message": "Validation failed"}), + get_routes={ + "/repos/langchain-ai/open-swe/branches/main": _FakeResponse(200, {"name": "main"}), + "/repos/langchain-ai/open-swe/branches/open-swe%2Ffeature": _FakeResponse( + 404, {"message": "Branch not found"} + ), + "/repos/langchain-ai/open-swe": _FakeResponse(200, {"private": True}), + }, + ) + _install_client(monkeypatch, client) + + result = _open() + + assert result["success"] is False + assert result["code"] == "github_pr_branch_not_visible" + assert result["branch_pushed"] is False + assert result["head_branch_visible"] is False + assert result["failed_step"] == "preflight_head_branch" + assert len(client.post_calls) == 1 + + async def _coro(value: Any) -> Any: return value @@ -355,7 +416,7 @@ def test_appends_plan_reference_from_thread_id(monkeypatch: pytest.MonkeyPatch) assert client.post_calls[0]["json"]["body"] == ( "body\n\n## References\n- Plan: https://dashboard.example/agents/thread-1/plan" ) - assert client.get_calls == [] + assert client.post_calls def test_omits_plan_reference_when_no_plan_exists(monkeypatch: pytest.MonkeyPatch) -> None: @@ -370,7 +431,7 @@ def test_omits_plan_reference_when_no_plan_exists(monkeypatch: pytest.MonkeyPatc _open_with_body("body") assert client.post_calls[0]["json"]["body"] == "body" - assert client.get_calls == [] + assert client.post_calls def test_omits_plan_reference_when_plan_markdown_empty(monkeypatch: pytest.MonkeyPatch) -> None: @@ -385,7 +446,7 @@ def test_omits_plan_reference_when_plan_markdown_empty(monkeypatch: pytest.Monke _open_with_body("body") assert client.post_calls[0]["json"]["body"] == "body" - assert client.get_calls == [] + assert client.post_calls def test_omits_plan_reference_when_store_lookup_fails(monkeypatch: pytest.MonkeyPatch) -> None: @@ -404,7 +465,7 @@ def test_omits_plan_reference_when_store_lookup_fails(monkeypatch: pytest.Monkey _open_with_body("body") assert client.post_calls[0]["json"]["body"] == "body" - assert client.get_calls == [] + assert client.post_calls def test_plan_reference_survives_source_reference_failure( @@ -433,7 +494,7 @@ def test_plan_reference_survives_source_reference_failure( sent_body = client.post_calls[0]["json"]["body"] assert "- Plan: https://dashboard.example/agents/thread-1/plan" in sent_body - assert client.get_calls == [] + assert client.post_calls def test_no_reference_for_public_repo(monkeypatch: pytest.MonkeyPatch) -> None: @@ -523,7 +584,7 @@ def test_skips_append_when_no_source_context(monkeypatch: pytest.MonkeyPatch) -> _open_with_body("body") assert client.post_calls[0]["json"]["body"] == "body" - assert client.get_calls == [] + assert client.post_calls def test_does_not_duplicate_existing_references(monkeypatch: pytest.MonkeyPatch) -> None: @@ -542,7 +603,7 @@ def test_does_not_duplicate_existing_references(monkeypatch: pytest.MonkeyPatch) _open_with_body("body\n\n## References\n- existing") assert client.post_calls[0]["json"]["body"] == "body\n\n## References\n- existing" - assert client.get_calls == [] + assert client.post_calls def test_derive_pr_state_prefers_merged() -> None: @@ -559,3 +620,44 @@ def test_derive_pr_state_draft() -> None: def test_derive_pr_state_open() -> None: assert opr.derive_pr_state(state="open", merged=False, draft=False) == "open" + + +def test_happy_path_skips_preflight_branch_gets(monkeypatch: pytest.MonkeyPatch) -> None: + """POST-first: when the PR is created successfully no preflight GETs + are made to /branches/... endpoints — only reference / plan checks.""" + _set_config(monkeypatch, {"source": "slack", "github_login": "johannes117"}) + monkeypatch.setattr(opr, "_resolve_pr_author_token", lambda: _coro(("tok", "user"))) + + client = _RoutingClient( + post=_FakeResponse(201, {"html_url": "u", "number": 1, "user": {}}), + get_routes={"/repos/langchain-ai/open-swe": _FakeResponse(200, {"private": False})}, + ) + _install_client(monkeypatch, client) + + result = _open() + + assert result["success"] is True + branch_gets = [c for c in client.get_calls if "/branches/" in c["url"]] + assert branch_gets == [] + + +def test_post_failure_diagnoses_via_preflight(monkeypatch: pytest.MonkeyPatch) -> None: + """When the POST fails with a non-201/non-422 status, the preflight runs + as a diagnostic and its result is returned (not the raw POST failure).""" + _set_config(monkeypatch, {"source": "slack", "github_login": "johannes117"}) + _stub_token(monkeypatch) + + client = _RoutingClient( + post=_FakeResponse(403, {"message": "Resource not accessible"}), + get_routes={ + "/repos/langchain-ai/open-swe": _FakeResponse(403, {"message": "Not Found"}), + }, + ) + _install_client(monkeypatch, client) + + result = _open() + + assert result["success"] is False + assert result["code"] == "github_app_access_missing_or_repo_not_found" + assert result["failed_step"] == "preflight_repo" + assert result["repo_visible"] is False diff --git a/tests/test_pr_creation_guard.py b/tests/test_pr_creation_guard.py new file mode 100644 index 00000000..55b9060d --- /dev/null +++ b/tests/test_pr_creation_guard.py @@ -0,0 +1,76 @@ +from __future__ import annotations + +import json +from typing import Any + +from langchain_core.messages import ToolMessage + +from agent.middleware.pr_creation_guard import ( + PullRequestCreationGuardMiddleware, + is_pr_creation_fallback_command, +) + + +class _Request: + def __init__(self, command: str) -> None: + self.tool_call = { + "name": "execute", + "args": {"command": command}, + "id": "call-1", + } + + +async def _handler(_request: Any) -> ToolMessage: + return ToolMessage(content="allowed", tool_call_id="call-1") + + +def test_detects_pr_creation_fallback_commands() -> None: + assert is_pr_creation_fallback_command("GH_TOKEN=dummy gh pr create --draft") + assert is_pr_creation_fallback_command( + "gh api repos/langchain-ai/open-swe/pulls -X POST -f title=x" + ) + assert is_pr_creation_fallback_command( + "gh api -X POST repos/langchain-ai/open-swe/pulls -f title=x" + ) + assert is_pr_creation_fallback_command( + "GH_TOKEN=dummy gh api -X POST repos/langchain-ai/open-swe/pulls -f title=x" + ) + assert is_pr_creation_fallback_command( + "curl -X POST https://api.github.com/repos/langchain-ai/open-swe/pulls -d '{}'" + ) + + +def test_allows_safe_pr_commands() -> None: + assert not is_pr_creation_fallback_command("GH_TOKEN=dummy gh pr view 1 --json url") + 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") + + +async def test_middleware_blocks_execute_pr_creation_fallbacks() -> None: + for command in ( + "GH_TOKEN=dummy gh pr create --draft", + "gh api repos/langchain-ai/open-swe/pulls -X POST -f title=x", + "GH_TOKEN=dummy gh api -X POST repos/langchain-ai/open-swe/pulls -f title=x", + "curl -X POST https://api.github.com/repos/langchain-ai/open-swe/pulls -d '{}'", + ): + result = await PullRequestCreationGuardMiddleware().awrap_tool_call( + _Request(command), _handler + ) + + assert isinstance(result, ToolMessage) + assert result.status == "error" + payload = json.loads(str(result.content)) + assert payload["code"] == "pr_creation_fallback_blocked" + assert payload["recoverable_by_agent"] is False + assert "open_pull_request" in payload["error"] + assert payload["blocked_command"] == command + + +async def test_middleware_allows_safe_pr_view() -> None: + result = await PullRequestCreationGuardMiddleware().awrap_tool_call( + _Request("GH_TOKEN=dummy gh pr view 1 --json url"), _handler + ) + + assert isinstance(result, ToolMessage) + assert result.content == "allowed"