mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-03 15:03:27 +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(
|
await create_github_pr(
|
||||||
repo_owner=repo_owner,
|
repo_owner=repo_owner,
|
||||||
repo_name=repo_name,
|
repo_name=repo_name,
|
||||||
github_token=installation_token,
|
github_token=github_token or installation_token,
|
||||||
title=pr_title,
|
title=pr_title,
|
||||||
head_branch=target_branch,
|
head_branch=target_branch,
|
||||||
base_branch=base_branch,
|
base_branch=base_branch,
|
||||||
body=pr_body,
|
body=pr_body,
|
||||||
|
installation_token=installation_token,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info("After-agent middleware completed successfully")
|
logger.info("After-agent middleware completed successfully")
|
||||||
|
|
|
||||||
|
|
@ -209,15 +209,17 @@ def commit_and_open_pr(
|
||||||
base_branch = asyncio.run(
|
base_branch = asyncio.run(
|
||||||
get_github_default_branch(repo_owner, repo_name, installation_token)
|
get_github_default_branch(repo_owner, repo_name, installation_token)
|
||||||
)
|
)
|
||||||
|
|
||||||
pr_url, _pr_number, pr_existing = asyncio.run(
|
pr_url, _pr_number, pr_existing = asyncio.run(
|
||||||
create_github_pr(
|
create_github_pr(
|
||||||
repo_owner=repo_owner,
|
repo_owner=repo_owner,
|
||||||
repo_name=repo_name,
|
repo_name=repo_name,
|
||||||
github_token=installation_token,
|
github_token=github_token or installation_token,
|
||||||
title=title,
|
title=title,
|
||||||
head_branch=target_branch,
|
head_branch=target_branch,
|
||||||
base_branch=base_branch,
|
base_branch=base_branch,
|
||||||
body=pr_body,
|
body=pr_body,
|
||||||
|
installation_token=installation_token,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -129,21 +129,35 @@ async def create_github_pr(
|
||||||
head_branch: str,
|
head_branch: str,
|
||||||
base_branch: str,
|
base_branch: str,
|
||||||
body: str,
|
body: str,
|
||||||
|
installation_token: str | None = None,
|
||||||
) -> tuple[str | None, int | None, bool]:
|
) -> tuple[str | None, int | None, bool]:
|
||||||
"""Create a draft GitHub pull request via the API.
|
"""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:
|
Args:
|
||||||
repo_owner: Repository owner (e.g., "langchain-ai")
|
repo_owner: Repository owner (e.g., "langchain-ai")
|
||||||
repo_name: Repository name (e.g., "deepagents")
|
repo_name: Repository name (e.g., "deepagents")
|
||||||
github_token: GitHub access token
|
github_token: GitHub access token (user token preferred)
|
||||||
title: PR title
|
title: PR title
|
||||||
head_branch: Source branch name
|
head_branch: Source branch name
|
||||||
base_branch: Target branch name
|
base_branch: Target branch name
|
||||||
body: PR description
|
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:
|
Returns:
|
||||||
Tuple of (pr_url, pr_number, pr_existing) if successful, (None, None, False) otherwise
|
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 = {
|
pr_payload = {
|
||||||
"title": title,
|
"title": title,
|
||||||
"head": head_branch,
|
"head": head_branch,
|
||||||
|
|
@ -161,52 +175,119 @@ async def create_github_pr(
|
||||||
)
|
)
|
||||||
|
|
||||||
async with httpx.AsyncClient() as http_client:
|
async with httpx.AsyncClient() as http_client:
|
||||||
try:
|
for token in tokens_to_try:
|
||||||
pr_response = await http_client.post(
|
try:
|
||||||
f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls",
|
pr_response = await http_client.post(
|
||||||
headers={
|
f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls",
|
||||||
"Authorization": f"Bearer {github_token}",
|
headers={
|
||||||
"Accept": "application/vnd.github+json",
|
"Authorization": f"Bearer {token}",
|
||||||
"X-GitHub-Api-Version": "2022-11-28",
|
"Accept": "application/vnd.github+json",
|
||||||
},
|
"X-GitHub-Api-Version": "2022-11-28",
|
||||||
json=pr_payload,
|
},
|
||||||
|
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,
|
||||||
)
|
)
|
||||||
|
except httpx.HTTPError:
|
||||||
pr_data = pr_response.json()
|
logger.warning("Failed to add label to PR #%s", pr_number, exc_info=True)
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
async def _find_existing_pr(
|
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