From 3405d145ac1eee1ea52bc2d4eaaf2326701b4839 Mon Sep 17 00:00:00 2001 From: Brace Sproul Date: Fri, 1 May 2026 18:12:41 -0700 Subject: [PATCH] feat: add edit_pull_request tool for editing PR titles/descriptions (#1063) * feat: add edit_pull_request tool for editing PR titles and descriptions * fix: patch auth flow in open PR middleware tests * fix: support app token for editing PRs --------- Co-authored-by: open-swe[bot] Co-authored-by: Johannes du Plessis --- agent/prompt.py | 3 + agent/server.py | 2 + agent/tools/__init__.py | 2 + agent/tools/edit_pull_request.py | 86 ++++++++++++++ agent/utils/github.py | 50 ++++++++ tests/test_edit_pull_request.py | 195 +++++++++++++++++++++++++++++++ tests/test_open_pr_middleware.py | 27 ++++- 7 files changed, 362 insertions(+), 3 deletions(-) create mode 100644 agent/tools/edit_pull_request.py create mode 100644 tests/test_edit_pull_request.py diff --git a/agent/prompt.py b/agent/prompt.py index cfde646a..946cdb1a 100644 --- a/agent/prompt.py +++ b/agent/prompt.py @@ -143,6 +143,9 @@ Do not use this tool to create or update the pull request for completed code cha #### `commit_and_open_pr` Commits all changes, pushes to a branch, and opens a **draft** GitHub PR. If a PR already exists for the branch, it is updated instead of recreated. +#### `edit_pull_request` +Edits the title and/or body of an existing GitHub Pull Request. Use this to update a PR description after creation — for example, after multiple iterations of changes. Requires `pr_number` and at least one of `title` or `body`. + #### `linear_comment` Posts a comment to a Linear ticket given a `ticket_id`. Call this **after** `commit_and_open_pr` to notify stakeholders that the work is done and include the PR link. You can tag Linear users with `@username` (their Linear display name). Example: "I've completed the implementation and opened a PR: . Hey @username, let me know if you have any feedback!". diff --git a/agent/server.py b/agent/server.py index 680c9e67..243efc2f 100644 --- a/agent/server.py +++ b/agent/server.py @@ -40,6 +40,7 @@ from .tools import ( commit_and_open_pr, create_pr_review, dismiss_pr_review, + edit_pull_request, fetch_url, get_branch_name, get_pr_check_runs, @@ -302,6 +303,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: list_repos, get_branch_name, commit_and_open_pr, + edit_pull_request, linear_comment, linear_create_issue, linear_delete_issue, diff --git a/agent/tools/__init__.py b/agent/tools/__init__.py index 07f38e32..07fd6fb2 100644 --- a/agent/tools/__init__.py +++ b/agent/tools/__init__.py @@ -1,4 +1,5 @@ from .commit_and_open_pr import commit_and_open_pr +from .edit_pull_request import edit_pull_request from .fetch_url import fetch_url from .get_branch_name import get_branch_name from .get_pr_review_comments import get_pr_review_comments @@ -30,6 +31,7 @@ __all__ = [ "commit_and_open_pr", "create_pr_review", "dismiss_pr_review", + "edit_pull_request", "fetch_url", "get_branch_name", "get_pr_check_runs", diff --git a/agent/tools/edit_pull_request.py b/agent/tools/edit_pull_request.py new file mode 100644 index 00000000..fc33062e --- /dev/null +++ b/agent/tools/edit_pull_request.py @@ -0,0 +1,86 @@ +import asyncio +import logging +from typing import Any + +from langgraph.config import get_config + +from ..utils.github import edit_github_pr +from ..utils.github_app import get_github_app_installation_token +from ..utils.github_token import get_github_token + +logger = logging.getLogger(__name__) + + +def edit_pull_request( + pr_number: int, + title: str | None = None, + body: str | None = None, +) -> dict[str, Any]: + """Edit the title and/or body of an existing GitHub Pull Request. + + Use this tool to update a PR's title or description after it has been created. + At least one of `title` or `body` must be provided. + + Args: + pr_number: The pull request number to edit. + title: New title for the PR. If not provided, the title is left unchanged. + body: New body/description for the PR. If not provided, the body is left unchanged. + + Returns: + Dictionary containing: + - success: Whether the operation completed successfully + - error: Error string if something failed, otherwise None + - pr_url: URL of the updated PR if successful, otherwise None + """ + try: + config = get_config() + configurable = config.get("configurable", {}) + + repo_config = configurable.get("repo", {}) + repo_owner = repo_config.get("owner") + repo_name = repo_config.get("name") + if not repo_owner or not repo_name: + return { + "success": False, + "error": "Missing repo owner/name in config", + "pr_url": None, + } + + if not pr_number: + return {"success": False, "error": "Missing pr_number argument", "pr_url": None} + + if title is None and body is None: + return { + "success": False, + "error": "At least one of title or body must be provided", + "pr_url": None, + } + + github_token = get_github_token(config) + if not github_token: + github_token = asyncio.run(get_github_app_installation_token()) + if not github_token: + return {"success": False, "error": "Missing GitHub token", "pr_url": None} + + pr_url, _pr_number = asyncio.run( + edit_github_pr( + repo_owner=repo_owner, + repo_name=repo_name, + github_token=github_token, + pr_number=pr_number, + title=title, + body=body, + ) + ) + + if not pr_url: + return { + "success": False, + "error": f"Failed to update PR #{pr_number}", + "pr_url": None, + } + + return {"success": True, "error": None, "pr_url": pr_url} + except Exception as e: + logger.exception("edit_pull_request failed") + return {"success": False, "error": f"{type(e).__name__}: {e}", "pr_url": None} diff --git a/agent/utils/github.py b/agent/utils/github.py index eae34dde..6a3f45f7 100644 --- a/agent/utils/github.py +++ b/agent/utils/github.py @@ -370,6 +370,56 @@ async def _find_existing_pr( return None, None +HTTP_OK = 200 + + +async def edit_github_pr( + repo_owner: str, + repo_name: str, + github_token: str, + pr_number: int, + title: str | None = None, + body: str | None = None, +) -> tuple[str | None, int | None]: + """Update an existing GitHub pull request title and/or body.""" + pr_payload: dict[str, str] = {} + if title is not None: + pr_payload["title"] = title + if body is not None: + pr_payload["body"] = body + + if not pr_payload: + logger.warning("edit_github_pr called with no fields to update") + return None, None + + async with httpx.AsyncClient() as http_client: + try: + response = await http_client.patch( + f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls/{pr_number}", + headers={ + "Authorization": f"Bearer {github_token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + }, + json=pr_payload, + ) + pr_data = response.json() + if response.status_code == HTTP_OK: + pr_url = pr_data.get("html_url") + logger.info("PR #%d updated successfully: %s", pr_number, pr_url) + return pr_url, pr_data.get("number") + + logger.error( + "GitHub API error (%s): %s", + response.status_code, + pr_data.get("message"), + ) + return None, None + except httpx.HTTPError: + logger.exception("Failed to update PR #%d via GitHub API", pr_number) + return None, None + + async def _update_github_pr( http_client: httpx.AsyncClient, repo_owner: str, diff --git a/tests/test_edit_pull_request.py b/tests/test_edit_pull_request.py new file mode 100644 index 00000000..7964ca39 --- /dev/null +++ b/tests/test_edit_pull_request.py @@ -0,0 +1,195 @@ +from __future__ import annotations + +import asyncio +import importlib +from typing import Any +from unittest.mock import AsyncMock + +import pytest + +from agent.tools import edit_pull_request as edit_pull_request_tool +from agent.utils import github + +edit_pull_request_module = importlib.import_module("agent.tools.edit_pull_request") + + +class _FakeResponse: + def __init__(self, status_code: int, payload: dict[str, Any]) -> None: + self.status_code = status_code + self._payload = payload + + def json(self) -> dict[str, Any]: + return self._payload + + +class _FakeAsyncClient: + def __init__( + self, + responses: list[_FakeResponse], + calls: list[tuple[str, str, dict[str, str], dict[str, str] | None]], + ) -> None: + self._responses = responses + self._calls = calls + + async def __aenter__(self) -> _FakeAsyncClient: + return self + + async def __aexit__(self, exc_type, exc, tb) -> None: + return None + + async def patch( + self, + url: str, + *, + headers: dict[str, str], + json: dict[str, str] | None = None, + ) -> _FakeResponse: + self._calls.append(("PATCH", url, headers, json)) + return self._responses.pop(0) + + +def _config() -> dict[str, Any]: + return {"configurable": {"repo": {"owner": "owner", "name": "repo"}}, "metadata": {}} + + +def test_edit_pull_request_requires_repo_config(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(edit_pull_request_module, "get_config", lambda: {"configurable": {}}) + + result = edit_pull_request_tool(pr_number=12, title="new title") + + assert result == { + "success": False, + "error": "Missing repo owner/name in config", + "pr_url": None, + } + + +def test_edit_pull_request_requires_update_field(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(edit_pull_request_module, "get_config", _config) + + result = edit_pull_request_tool(pr_number=12) + + assert result == { + "success": False, + "error": "At least one of title or body must be provided", + "pr_url": None, + } + + +def test_edit_pull_request_prefers_user_token(monkeypatch: pytest.MonkeyPatch) -> None: + edit_mock = AsyncMock(return_value=("https://github.com/owner/repo/pull/12", 12)) + app_token_mock = AsyncMock(return_value="app-token") + monkeypatch.setattr(edit_pull_request_module, "get_config", _config) + monkeypatch.setattr(edit_pull_request_module, "get_github_token", lambda config: "user-token") + monkeypatch.setattr( + edit_pull_request_module, + "get_github_app_installation_token", + app_token_mock, + ) + monkeypatch.setattr(edit_pull_request_module, "edit_github_pr", edit_mock) + + result = edit_pull_request_tool(pr_number=12, title="new title") + + assert result == { + "success": True, + "error": None, + "pr_url": "https://github.com/owner/repo/pull/12", + } + app_token_mock.assert_not_called() + edit_mock.assert_awaited_once_with( + repo_owner="owner", + repo_name="repo", + github_token="user-token", + pr_number=12, + title="new title", + body=None, + ) + + +def test_edit_pull_request_falls_back_to_app_token(monkeypatch: pytest.MonkeyPatch) -> None: + edit_mock = AsyncMock(return_value=("https://github.com/owner/repo/pull/12", 12)) + monkeypatch.setattr(edit_pull_request_module, "get_config", _config) + monkeypatch.setattr(edit_pull_request_module, "get_github_token", lambda config: None) + monkeypatch.setattr( + edit_pull_request_module, + "get_github_app_installation_token", + AsyncMock(return_value="app-token"), + ) + monkeypatch.setattr(edit_pull_request_module, "edit_github_pr", edit_mock) + + result = edit_pull_request_tool(pr_number=12, body="new body") + + assert result["success"] is True + edit_mock.assert_awaited_once_with( + repo_owner="owner", + repo_name="repo", + github_token="app-token", + pr_number=12, + title=None, + body="new body", + ) + + +def test_edit_pull_request_fails_without_any_token(monkeypatch: pytest.MonkeyPatch) -> None: + edit_mock = AsyncMock() + monkeypatch.setattr(edit_pull_request_module, "get_config", _config) + monkeypatch.setattr(edit_pull_request_module, "get_github_token", lambda config: None) + monkeypatch.setattr( + edit_pull_request_module, + "get_github_app_installation_token", + AsyncMock(return_value=None), + ) + monkeypatch.setattr(edit_pull_request_module, "edit_github_pr", edit_mock) + + result = edit_pull_request_tool(pr_number=12, title="new title") + + assert result == {"success": False, "error": "Missing GitHub token", "pr_url": None} + edit_mock.assert_not_called() + + +def test_edit_github_pr_sends_partial_patch_payload(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[str, str, dict[str, str], dict[str, str] | None]] = [] + responses = [_FakeResponse(200, {"html_url": "https://github.com/o/r/pull/12", "number": 12})] + monkeypatch.setattr(github.httpx, "AsyncClient", lambda: _FakeAsyncClient(responses, calls)) + + result = asyncio.run( + github.edit_github_pr( + repo_owner="o", + repo_name="r", + github_token="token", + pr_number=12, + title="new title", + ) + ) + + assert result == ("https://github.com/o/r/pull/12", 12) + assert calls == [ + ( + "PATCH", + "https://api.github.com/repos/o/r/pulls/12", + { + "Authorization": "Bearer token", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + }, + {"title": "new title"}, + ) + ] + + +def test_edit_github_pr_returns_none_on_api_failure(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[str, str, dict[str, str], dict[str, str] | None]] = [] + responses = [_FakeResponse(404, {"message": "Not Found"})] + monkeypatch.setattr(github.httpx, "AsyncClient", lambda: _FakeAsyncClient(responses, calls)) + + result = asyncio.run( + github.edit_github_pr( + repo_owner="o", + repo_name="r", + github_token="token", + pr_number=12, + body="new body", + ) + ) + + assert result == (None, None) diff --git a/tests/test_open_pr_middleware.py b/tests/test_open_pr_middleware.py index 8bfe3a4f..5dbdc73c 100644 --- a/tests/test_open_pr_middleware.py +++ b/tests/test_open_pr_middleware.py @@ -5,6 +5,7 @@ the success value from commit_and_open_pr tool results. """ import json +from contextlib import ExitStack from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -93,6 +94,23 @@ class TestOpenPrIfNeededMiddleware: def _make_state(self, messages: list) -> dict: return {"messages": messages} + def _patch_auth_flow(self) -> ExitStack: + stack = ExitStack() + stack.enter_context( + patch("agent.middleware.open_pr.get_github_token", return_value="token") + ) + stack.enter_context( + patch("agent.middleware.open_pr.resolve_triggering_user_identity", return_value=None) + ) + stack.enter_context( + patch( + "agent.middleware.open_pr.get_github_app_installation_token", + new_callable=AsyncMock, + return_value="installation-token", + ) + ) + return stack + def test_skips_when_commit_and_open_pr_succeeded(self) -> None: """When success=True, the tool handled everything — middleware should be a no-op.""" payload = {"success": True, "error": None, "pr_url": "https://github.com/org/repo/pull/42"} @@ -239,7 +257,8 @@ class TestOpenPrIfNeededMiddleware: ): # Middleware should NOT short-circuit; it reaches sandbox logic # We verify get_sandbox_backend was called (safety net fired) - await open_pr_if_needed.aafter_agent(state, self._make_runtime()) + with self._patch_auth_flow(): + await open_pr_if_needed.aafter_agent(state, self._make_runtime()) # The safety net fired: get_sandbox_backend was called mock_sandbox.assert_called_once_with("thread-push-fail") @@ -289,7 +308,8 @@ class TestOpenPrIfNeededMiddleware: "agent.middleware.open_pr.git_has_unpushed_commits", return_value=True, ): - await open_pr_if_needed.aafter_agent(state, self._make_runtime()) + with self._patch_auth_flow(): + await open_pr_if_needed.aafter_agent(state, self._make_runtime()) # The safety net fired: get_sandbox_backend was called mock_sandbox.assert_called_once_with("thread-pr-fail") @@ -353,7 +373,8 @@ class TestOpenPrIfNeededMiddleware: with patch( "agent.middleware.open_pr.get_sandbox_backend", side_effect=fake_get_sandbox ): - await open_pr_if_needed.aafter_agent(state, self._make_runtime()) + with self._patch_auth_flow(): + await open_pr_if_needed.aafter_agent(state, self._make_runtime()) # If the old buggy `"success" in pr_payload` check was used, the middleware # would have returned None before reaching get_sandbox_backend.