diff --git a/agent/utils/github.py b/agent/utils/github.py index 850ccfc1..18285656 100644 --- a/agent/utils/github.py +++ b/agent/utils/github.py @@ -222,16 +222,31 @@ async def create_github_pr( github_token=token, head_branch=head_branch, ) - if existing: + pr_url, pr_number = existing + if pr_url: + logger.info("Using existing PR for head branch: %s", pr_url) + updated = await _update_github_pr( + http_client=http_client, + repo_owner=repo_owner, + repo_name=repo_name, + github_token=token, + pr_number=pr_number, + title=title, + body=body, + ) + if not updated: + if token != tokens_to_try[-1]: + logger.info("Retrying existing PR update with installation token") + continue + return None, None, False await _add_label( http_client, repo_owner, repo_name, label_tok, - existing[1], + pr_number, ) - logger.info("Using existing PR for head branch: %s", existing[0]) - return existing[0], existing[1], True + return pr_url, pr_number, True else: logger.debug( "Could not find existing PR with current token, will retry" @@ -353,6 +368,45 @@ async def _find_existing_pr( return None, None +async def _update_github_pr( + http_client: httpx.AsyncClient, + repo_owner: str, + repo_name: str, + github_token: str, + pr_number: int | None, + title: str, + body: str, +) -> bool: + """Update an existing PR's title and body via PATCH.""" + if pr_number is None: + logger.warning("Cannot update PR: pr_number is None") + return False + headers = { + "Authorization": f"Bearer {github_token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + } + try: + response = await http_client.patch( + f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls/{pr_number}", + headers=headers, + json={"title": title, "body": body}, + ) + except httpx.HTTPError: + logger.warning("Failed to update PR #%s", pr_number, exc_info=True) + return False + if response.status_code == 200: # noqa: PLR2004 + logger.info("Updated existing PR #%s with new title and body", pr_number) + return True + logger.warning( + "Failed to update PR #%s (%s): %s", + pr_number, + response.status_code, + response.json().get("message"), + ) + return False + + async def get_github_default_branch( repo_owner: str, repo_name: str, diff --git a/tests/test_github_pr_label.py b/tests/test_github_pr_label.py index 34b90b68..3274fc2f 100644 --- a/tests/test_github_pr_label.py +++ b/tests/test_github_pr_label.py @@ -46,6 +46,12 @@ class _FakeAsyncClient: self._calls.append(("GET", url, params)) return self._responses.pop(0) + async def patch( + self, url: str, *, headers: dict[str, str], json: dict | None = None + ) -> _FakeResponse: + self._calls.append(("PATCH", url, json)) + return self._responses.pop(0) + class _RaiseOnLabelPostClient(_FakeAsyncClient): """Raises on the second POST (the label call) to simulate network failure.""" @@ -161,6 +167,7 @@ def test_create_pr_adds_label_on_existing_pr(monkeypatch: pytest.MonkeyPatch) -> responses = [ _FakeResponse(422, {"message": "A pull request already exists"}), _FakeResponse(200, [{"html_url": "https://github.com/o/r/pull/7", "number": 7}]), + _FakeResponse(200, {"html_url": "https://github.com/o/r/pull/7", "number": 7}), _FakeResponse(200, [{"name": "OpenSWE"}]), ] monkeypatch.setattr(github.httpx, "AsyncClient", lambda: _FakeAsyncClient(responses, calls)) @@ -180,12 +187,98 @@ def test_create_pr_adds_label_on_existing_pr(monkeypatch: pytest.MonkeyPatch) -> assert result == ("https://github.com/o/r/pull/7", 7, True) assert calls[2] == ( + "PATCH", + "https://api.github.com/repos/o/r/pulls/7", + {"title": "feat: test", "body": "body"}, + ) + assert calls[3] == ( "POST", "https://api.github.com/repos/o/r/issues/7/labels", {"labels": ["OpenSWE"]}, ) +def test_create_pr_returns_failure_when_existing_pr_update_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + _FakeResponse(422, {"message": "A pull request already exists"}), + _FakeResponse(200, [{"html_url": "https://github.com/o/r/pull/7", "number": 7}]), + _FakeResponse(403, {"message": "Resource not accessible by integration"}), + ] + monkeypatch.setattr(github.httpx, "AsyncClient", lambda: _FakeAsyncClient(responses, calls)) + + result = asyncio.run( + github.create_github_pr( + repo_owner="o", + repo_name="r", + github_token="token", + title="feat: test", + head_branch="feature", + base_branch="main", + body="body", + ) + ) + + assert result == (None, None, False) + assert calls == [ + ( + "POST", + "https://api.github.com/repos/o/r/pulls", + { + "title": "feat: test", + "head": "feature", + "base": "main", + "body": "body", + "draft": True, + }, + ), + ( + "GET", + "https://api.github.com/repos/o/r/pulls", + {"head": "o:feature", "state": "open", "per_page": 1}, + ), + ( + "PATCH", + "https://api.github.com/repos/o/r/pulls/7", + {"title": "feat: test", "body": "body"}, + ), + ] + + +def test_create_pr_retries_existing_pr_update_with_installation_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + _FakeResponse(422, {"message": "A pull request already exists"}), + _FakeResponse(200, [{"html_url": "https://github.com/o/r/pull/7", "number": 7}]), + _FakeResponse(403, {"message": "Resource not accessible by integration"}), + _FakeResponse(422, {"message": "A pull request already exists"}), + _FakeResponse(200, [{"html_url": "https://github.com/o/r/pull/7", "number": 7}]), + _FakeResponse(200, {"html_url": "https://github.com/o/r/pull/7", "number": 7}), + _FakeResponse(200, [{"name": "OpenSWE"}]), + ] + monkeypatch.setattr(github.httpx, "AsyncClient", lambda: _FakeAsyncClient(responses, calls)) + + result = asyncio.run( + github.create_github_pr( + repo_owner="o", + repo_name="r", + github_token="user-token", + title="feat: test", + head_branch="feature", + base_branch="main", + body="body", + installation_token="install-token", + ) + ) + + assert result == ("https://github.com/o/r/pull/7", 7, True) + assert [call[0] for call in calls] == ["POST", "GET", "PATCH", "POST", "GET", "PATCH", "POST"] + + def test_create_pr_succeeds_when_label_fails(monkeypatch: pytest.MonkeyPatch) -> None: """PR creation should succeed even if labeling fails.""" calls: list[tuple[str, str, dict | None]] = []