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] <open-swe@users.noreply.github.com>
This commit is contained in:
Aran Yogesh 2026-04-29 17:34:00 -07:00 • committed by GitHub
parent 448be4a466
commit 920a8a7624
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 481 additions and 48 deletions

View file

@ -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")

View file

@ -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,
)
)

View file

@ -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(

View file

@ -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