mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 05:43:14 +00:00
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] <open-swe@users.noreply.github.com> Co-authored-by: Johannes du Plessis <johannes@langchain.dev>
This commit is contained in:
parent
3d3d8403fe
commit
3405d145ac
7 changed files with 362 additions and 3 deletions
|
|
@ -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: <pr_url>. Hey @username, let me know if you have any feedback!".
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
86
agent/tools/edit_pull_request.py
Normal file
86
agent/tools/edit_pull_request.py
Normal file
|
|
@ -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}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
195
tests/test_edit_pull_request.py
Normal file
195
tests/test_edit_pull_request.py
Normal file
|
|
@ -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)
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue