From 920a8a7624d09e88cc6f9e361898dc2d4aba5707 Mon Sep 17 00:00:00 2001 From: Aran Yogesh Date: Wed, 29 Apr 2026 17:34:00 -0700 Subject: [PATCH] feat: open PRs under user's name and add OpenSWE label (#1215) * feat: open PRs under user's name and add OpenSWE label * feat: use user token for PR authorship, add OpenSWE label, and consolidate fallback logic * linting * fix: address review nits for PR authorship and labeling Fix docstring casing, add debug logging for 422 existing-PR search fallback, tighten test type annotations, and add missing HTTPError fallback test. --------- Co-authored-by: open-swe[bot] --- agent/middleware/open_pr.py | 3 +- agent/tools/commit_and_open_pr.py | 4 +- agent/utils/github.py | 173 +++++++++++---- tests/test_github_pr_label.py | 349 ++++++++++++++++++++++++++++++ 4 files changed, 481 insertions(+), 48 deletions(-) create mode 100644 tests/test_github_pr_label.py diff --git a/agent/middleware/open_pr.py b/agent/middleware/open_pr.py index 12d9325e..1b6404ce 100644 --- a/agent/middleware/open_pr.py +++ b/agent/middleware/open_pr.py @@ -168,11 +168,12 @@ async def open_pr_if_needed( await create_github_pr( repo_owner=repo_owner, repo_name=repo_name, - github_token=installation_token, + github_token=github_token or installation_token, title=pr_title, head_branch=target_branch, base_branch=base_branch, body=pr_body, + installation_token=installation_token, ) logger.info("After-agent middleware completed successfully") diff --git a/agent/tools/commit_and_open_pr.py b/agent/tools/commit_and_open_pr.py index b4922c91..cd8e277e 100644 --- a/agent/tools/commit_and_open_pr.py +++ b/agent/tools/commit_and_open_pr.py @@ -209,15 +209,17 @@ def commit_and_open_pr( base_branch = asyncio.run( get_github_default_branch(repo_owner, repo_name, installation_token) ) + pr_url, _pr_number, pr_existing = asyncio.run( create_github_pr( repo_owner=repo_owner, repo_name=repo_name, - github_token=installation_token, + github_token=github_token or installation_token, title=title, head_branch=target_branch, base_branch=base_branch, body=pr_body, + installation_token=installation_token, ) ) diff --git a/agent/utils/github.py b/agent/utils/github.py index be7fc835..643883c5 100644 --- a/agent/utils/github.py +++ b/agent/utils/github.py @@ -129,21 +129,35 @@ async def create_github_pr( head_branch: str, base_branch: str, body: str, + installation_token: str | None = None, ) -> tuple[str | None, int | None, bool]: """Create a draft GitHub pull request via the API. + When *github_token* differs from *installation_token* (e.g. a user + OAuth token), the function first attempts to create the PR with the + user token so the user becomes the PR author. If that fails it + retries with the installation token. The ``OpenSWE`` label is + always added using the installation token. + Args: repo_owner: Repository owner (e.g., "langchain-ai") repo_name: Repository name (e.g., "deepagents") - github_token: GitHub access token + github_token: GitHub access token (user token preferred) title: PR title head_branch: Source branch name base_branch: Target branch name body: PR description + installation_token: GitHub App installation token used for labeling and as a fallback + for PR creation. Falls back to github_token when not provided. Returns: Tuple of (pr_url, pr_number, pr_existing) if successful, (None, None, False) otherwise """ + tokens_to_try = [github_token] + if installation_token and installation_token != github_token: + tokens_to_try.append(installation_token) + label_tok = installation_token or github_token + pr_payload = { "title": title, "head": head_branch, @@ -161,52 +175,119 @@ async def create_github_pr( ) async with httpx.AsyncClient() as http_client: - try: - pr_response = await http_client.post( - f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls", - headers={ - "Authorization": f"Bearer {github_token}", - "Accept": "application/vnd.github+json", - "X-GitHub-Api-Version": "2022-11-28", - }, - json=pr_payload, + for token in tokens_to_try: + try: + pr_response = await http_client.post( + f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls", + headers={ + "Authorization": f"Bearer {token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + }, + json=pr_payload, + ) + + pr_data = pr_response.json() + + if pr_response.status_code == HTTP_CREATED: + pr_url = pr_data.get("html_url") + pr_number = pr_data.get("number") + await _add_label( + http_client, + repo_owner, + repo_name, + label_tok, + pr_number, + ) + logger.info("PR created successfully: %s", pr_url) + return pr_url, pr_number, False + + if pr_response.status_code == HTTP_UNPROCESSABLE_ENTITY: + logger.error("GitHub API validation error (422): %s", pr_data.get("message")) + existing = await _find_existing_pr( + http_client=http_client, + repo_owner=repo_owner, + repo_name=repo_name, + github_token=token, + head_branch=head_branch, + ) + if existing: + await _add_label( + http_client, + repo_owner, + repo_name, + label_tok, + existing[1], + ) + logger.info("Using existing PR for head branch: %s", existing[0]) + return existing[0], existing[1], True + else: + logger.debug( + "Could not find existing PR with current token, will retry" + if token != tokens_to_try[-1] + else "Could not find existing PR" + ) + else: + logger.error( + "GitHub API error (%s): %s", + pr_response.status_code, + pr_data.get("message"), + ) + + if "errors" in pr_data: + logger.error("GitHub API errors detail: %s", pr_data.get("errors")) + + # If this was the user token, fall through to retry with installation token + if token != tokens_to_try[-1]: + logger.info("Retrying PR creation with installation token") + continue + + return None, None, False + + except httpx.HTTPError: + logger.exception("Failed to create PR via GitHub API") + if token != tokens_to_try[-1]: + logger.info("Retrying PR creation with installation token") + continue + return None, None, False + + return None, None, False + + +_OPENSWE_LABEL = "OpenSWE" + + +async def _add_label( + http_client: httpx.AsyncClient, + repo_owner: str, + repo_name: str, + github_token: str, + pr_number: int | None, +) -> None: + """Add the 'OpenSWE' label to a PR without failing PR creation on errors.""" + if not pr_number: + return + + try: + response = await http_client.post( + f"https://api.github.com/repos/{repo_owner}/{repo_name}/issues/{pr_number}/labels", + headers={ + "Authorization": f"Bearer {github_token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + }, + json={"labels": [_OPENSWE_LABEL]}, + ) + if response.is_success: + logger.info("Added '%s' label to PR #%s", _OPENSWE_LABEL, pr_number) + else: + logger.warning( + "Failed to add label to PR #%s (%s)", + pr_number, + response.status_code, ) - - pr_data = pr_response.json() - - if pr_response.status_code == HTTP_CREATED: - pr_url = pr_data.get("html_url") - pr_number = pr_data.get("number") - logger.info("PR created successfully: %s", pr_url) - return pr_url, pr_number, False - - if pr_response.status_code == HTTP_UNPROCESSABLE_ENTITY: - logger.error("GitHub API validation error (422): %s", pr_data.get("message")) - existing = await _find_existing_pr( - http_client=http_client, - repo_owner=repo_owner, - repo_name=repo_name, - github_token=github_token, - head_branch=head_branch, - ) - if existing: - logger.info("Using existing PR for head branch: %s", existing[0]) - return existing[0], existing[1], True - else: - logger.error( - "GitHub API error (%s): %s", - pr_response.status_code, - pr_data.get("message"), - ) - - if "errors" in pr_data: - logger.error("GitHub API errors detail: %s", pr_data.get("errors")) - - return None, None, False - - except httpx.HTTPError: - logger.exception("Failed to create PR via GitHub API") - return None, None, False + except httpx.HTTPError: + logger.warning("Failed to add label to PR #%s", pr_number, exc_info=True) async def _find_existing_pr( diff --git a/tests/test_github_pr_label.py b/tests/test_github_pr_label.py new file mode 100644 index 00000000..34b90b68 --- /dev/null +++ b/tests/test_github_pr_label.py @@ -0,0 +1,349 @@ +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from agent.utils import github + + +class _FakeResponse: + def __init__(self, status_code: int, payload: Any) -> None: + self.status_code = status_code + self._payload = payload + + @property + def is_success(self) -> bool: + return 200 <= self.status_code < 300 + + def json(self) -> Any: + return self._payload + + +class _FakeAsyncClient: + def __init__( + self, responses: list[_FakeResponse], calls: list[tuple[str, str, dict | 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 post( + self, url: str, *, headers: dict[str, str], json: dict | None = None + ) -> _FakeResponse: + self._calls.append(("POST", url, json)) + return self._responses.pop(0) + + async def get( + self, url: str, *, headers: dict[str, str], params: dict[str, str | int] + ) -> _FakeResponse: + self._calls.append(("GET", url, params)) + return self._responses.pop(0) + + +class _RaiseOnLabelPostClient(_FakeAsyncClient): + """Raises on the second POST (the label call) to simulate network failure.""" + + async def post( + self, url: str, *, headers: dict[str, str], json: dict | None = None + ) -> _FakeResponse: + self._calls.append(("POST", url, json)) + if len(self._calls) == 2: + request = github.httpx.Request("POST", url) + raise github.httpx.ConnectError("boom", request=request) + return self._responses.pop(0) + + +class _AlwaysRaisePostClient(_FakeAsyncClient): + """Raises on every POST to simulate network failure.""" + + async def post( + self, url: str, *, headers: dict[str, str], json: dict | None = None + ) -> _FakeResponse: + self._calls.append(("POST", url, json)) + request = github.httpx.Request("POST", url) + raise github.httpx.ConnectError("boom", request=request) + + +# -- _add_label tests -- + + +def test_add_label_success(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[str, str, dict | None]] = [] + responses = [_FakeResponse(200, [{"name": "OpenSWE"}])] + client = _FakeAsyncClient(responses, calls) + + asyncio.run(github._add_label(client, "o", "r", "token", 12)) + + assert calls == [ + ("POST", "https://api.github.com/repos/o/r/issues/12/labels", {"labels": ["OpenSWE"]}), + ] + + +def test_add_label_skips_when_no_pr_number() -> None: + """Should return immediately without making any API calls.""" + + async def _run() -> None: + # Pass a mock that would fail if called + await github._add_label(None, "o", "r", "token", None) # type: ignore[arg-type] + + asyncio.run(_run()) + + +def test_add_label_does_not_raise_on_api_failure(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[str, str, dict | None]] = [] + responses = [_FakeResponse(403, {"message": "Resource not accessible by integration"})] + client = _FakeAsyncClient(responses, calls) + + # Should not raise + asyncio.run(github._add_label(client, "o", "r", "token", 12)) + + assert len(calls) == 1 + + +def test_add_label_does_not_raise_on_http_error() -> None: + calls: list[tuple[str, str, dict | None]] = [] + responses: list[_FakeResponse] = [] + client = _AlwaysRaisePostClient(responses, calls) + + # The POST will raise — should not propagate + asyncio.run(github._add_label(client, "o", "r", "token", 5)) + + assert len(calls) == 1 + + +# -- create_github_pr with label tests -- + + +def test_create_pr_adds_label_on_new_pr(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + _FakeResponse(201, {"html_url": "https://github.com/o/r/pull/12", "number": 12}), + _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/12", 12, False) + # First call: create PR with user token, second call: add label with install token + assert calls[0] == ( + "POST", + "https://api.github.com/repos/o/r/pulls", + {"title": "feat: test", "head": "feature", "base": "main", "body": "body", "draft": True}, + ) + assert calls[1] == ( + "POST", + "https://api.github.com/repos/o/r/issues/12/labels", + {"labels": ["OpenSWE"]}, + ) + + +def test_create_pr_adds_label_on_existing_pr(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(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 calls[2] == ( + "POST", + "https://api.github.com/repos/o/r/issues/7/labels", + {"labels": ["OpenSWE"]}, + ) + + +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]] = [] + responses = [ + _FakeResponse(201, {"html_url": "https://github.com/o/r/pull/12", "number": 12}), + ] + monkeypatch.setattr( + github.httpx, "AsyncClient", lambda: _RaiseOnLabelPostClient(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 == ("https://github.com/o/r/pull/12", 12, False) + + +def test_create_pr_uses_github_token_for_label_when_no_installation_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """When installation_token is not provided, github_token is used for labeling.""" + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + _FakeResponse(201, {"html_url": "https://github.com/o/r/pull/12", "number": 12}), + _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="token", + title="feat: test", + head_branch="feature", + base_branch="main", + body="body", + ) + ) + + assert result == ("https://github.com/o/r/pull/12", 12, False) + # Both calls made — label uses the same token + assert len(calls) == 2 + + +def test_create_pr_falls_back_to_installation_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """When user token fails, retries with installation token.""" + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + # First attempt with user token → 403 + _FakeResponse(403, {"message": "Resource not accessible by integration"}), + # Second attempt with installation token → 201 + _FakeResponse(201, {"html_url": "https://github.com/o/r/pull/12", "number": 12}), + # Label + _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/12", 12, False) + # 3 calls: failed PR create, successful PR create, label + assert len(calls) == 3 + + +def test_create_pr_falls_back_on_http_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """When user token raises HTTPError, retries with installation token.""" + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + # Installation token succeeds + _FakeResponse(201, {"html_url": "https://github.com/o/r/pull/12", "number": 12}), + # Label + _FakeResponse(200, [{"name": "OpenSWE"}]), + ] + + class _RaiseFirstPostClient(_FakeAsyncClient): + """Raises on the first POST only (user token), then delegates to normal behavior.""" + + _first = True + + async def post( + self, url: str, *, headers: dict[str, str], json: dict | None = None + ) -> _FakeResponse: + self._calls.append(("POST", url, json)) + if self._first: + self._first = False + request = github.httpx.Request("POST", url) + raise github.httpx.ConnectError("boom", request=request) + return self._responses.pop(0) + + monkeypatch.setattr( + github.httpx, "AsyncClient", lambda: _RaiseFirstPostClient(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/12", 12, False) + # 3 calls: failed POST (user token), successful POST (install token), label POST + assert len(calls) == 3 + + +def test_create_pr_no_fallback_when_tokens_are_same( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """When github_token == installation_token, no retry happens.""" + calls: list[tuple[str, str, dict | None]] = [] + responses = [ + _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="same-token", + title="feat: test", + head_branch="feature", + base_branch="main", + body="body", + installation_token="same-token", + ) + ) + + assert result == (None, None, False) + # Only 1 call — no retry since tokens are identical + assert len(calls) == 1