mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 08:03:15 +00:00
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:
parent
448be4a466
commit
920a8a7624
4 changed files with 481 additions and 48 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
349
tests/test_github_pr_label.py
Normal file
349
tests/test_github_pr_label.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue