mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 10:23:14 +00:00
fix: Auto assign PRs to creator (#1211)
* fix: Auto assign PRs to creator * cr
This commit is contained in:
parent
5925a90a95
commit
e9b94ac8ba
5 changed files with 282 additions and 0 deletions
|
|
@ -173,6 +173,7 @@ async def open_pr_if_needed(
|
|||
head_branch=target_branch,
|
||||
base_branch=base_branch,
|
||||
body=pr_body,
|
||||
assignee_login=user_identity.github_login if user_identity else None,
|
||||
)
|
||||
|
||||
logger.info("After-agent middleware completed successfully")
|
||||
|
|
|
|||
|
|
@ -218,6 +218,7 @@ def commit_and_open_pr(
|
|||
head_branch=target_branch,
|
||||
base_branch=base_branch,
|
||||
body=pr_body,
|
||||
assignee_login=user_identity.github_login if user_identity else None,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ class CollaboratorIdentity:
|
|||
display_name: str
|
||||
commit_name: str
|
||||
commit_email: str
|
||||
github_login: str | None = None
|
||||
|
||||
|
||||
def _normalize_text(value: Any) -> str:
|
||||
|
|
@ -72,6 +73,7 @@ def _identity_from_github_token(github_token: str | None) -> CollaboratorIdentit
|
|||
display_name=display_name,
|
||||
commit_name=display_name,
|
||||
commit_email=commit_email,
|
||||
github_login=login,
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
logger.debug("Failed to resolve GitHub user identity from token", exc_info=True)
|
||||
|
|
@ -92,6 +94,7 @@ def _identity_from_config(config: dict[str, Any]) -> CollaboratorIdentity | None
|
|||
display_name=github_login,
|
||||
commit_name=github_login,
|
||||
commit_email=commit_email,
|
||||
github_login=github_login,
|
||||
)
|
||||
|
||||
slack_thread = configurable.get("slack_thread", {})
|
||||
|
|
|
|||
|
|
@ -129,6 +129,7 @@ async def create_github_pr(
|
|||
head_branch: str,
|
||||
base_branch: str,
|
||||
body: str,
|
||||
assignee_login: str | None = None,
|
||||
) -> tuple[str | None, int | None, bool]:
|
||||
"""Create a draft GitHub pull request via the API.
|
||||
|
||||
|
|
@ -140,6 +141,7 @@ async def create_github_pr(
|
|||
head_branch: Source branch name
|
||||
base_branch: Target branch name
|
||||
body: PR description
|
||||
assignee_login: GitHub login to assign to the PR after it is opened
|
||||
|
||||
Returns:
|
||||
Tuple of (pr_url, pr_number, pr_existing) if successful, (None, None, False) otherwise
|
||||
|
|
@ -177,6 +179,14 @@ async def create_github_pr(
|
|||
if pr_response.status_code == HTTP_CREATED:
|
||||
pr_url = pr_data.get("html_url")
|
||||
pr_number = pr_data.get("number")
|
||||
await _assign_pr_assignee(
|
||||
http_client=http_client,
|
||||
repo_owner=repo_owner,
|
||||
repo_name=repo_name,
|
||||
github_token=github_token,
|
||||
pr_number=pr_number,
|
||||
assignee_login=assignee_login,
|
||||
)
|
||||
logger.info("PR created successfully: %s", pr_url)
|
||||
return pr_url, pr_number, False
|
||||
|
||||
|
|
@ -190,6 +200,14 @@ async def create_github_pr(
|
|||
head_branch=head_branch,
|
||||
)
|
||||
if existing:
|
||||
await _assign_pr_assignee(
|
||||
http_client=http_client,
|
||||
repo_owner=repo_owner,
|
||||
repo_name=repo_name,
|
||||
github_token=github_token,
|
||||
pr_number=existing[1],
|
||||
assignee_login=assignee_login,
|
||||
)
|
||||
logger.info("Using existing PR for head branch: %s", existing[0])
|
||||
return existing[0], existing[1], True
|
||||
else:
|
||||
|
|
@ -209,6 +227,52 @@ async def create_github_pr(
|
|||
return None, None, False
|
||||
|
||||
|
||||
async def _assign_pr_assignee(
|
||||
http_client: httpx.AsyncClient,
|
||||
repo_owner: str,
|
||||
repo_name: str,
|
||||
github_token: str,
|
||||
pr_number: int | None,
|
||||
assignee_login: str | None,
|
||||
) -> None:
|
||||
"""Assign a PR to a GitHub user without failing PR creation on errors."""
|
||||
if not pr_number or not assignee_login:
|
||||
return
|
||||
|
||||
try:
|
||||
response = await http_client.post(
|
||||
f"https://api.github.com/repos/{repo_owner}/{repo_name}/issues/{pr_number}/assignees",
|
||||
headers={
|
||||
"Authorization": f"Bearer {github_token}",
|
||||
"Accept": "application/vnd.github+json",
|
||||
"X-GitHub-Api-Version": "2022-11-28",
|
||||
},
|
||||
json={"assignees": [assignee_login]},
|
||||
)
|
||||
if response.is_success:
|
||||
logger.info("Assigned PR #%s to %s", pr_number, assignee_login)
|
||||
return
|
||||
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError:
|
||||
payload = {}
|
||||
logger.warning(
|
||||
"Failed to assign PR #%s to %s (%s): %s",
|
||||
pr_number,
|
||||
assignee_login,
|
||||
response.status_code,
|
||||
payload.get("message") if isinstance(payload, dict) else None,
|
||||
)
|
||||
except httpx.HTTPError:
|
||||
logger.warning(
|
||||
"Failed to assign PR #%s to %s because of an HTTP error",
|
||||
pr_number,
|
||||
assignee_login,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
async def _find_existing_pr(
|
||||
http_client: httpx.AsyncClient,
|
||||
repo_owner: str,
|
||||
|
|
|
|||
213
tests/test_github_pr_assignment.py
Normal file
213
tests/test_github_pr_assignment.py
Normal file
|
|
@ -0,0 +1,213 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from agent.utils import authorship, github
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, status_code: int, payload: object) -> None:
|
||||
self.status_code = status_code
|
||||
self._payload = payload
|
||||
|
||||
@property
|
||||
def is_success(self) -> bool:
|
||||
return 200 <= self.status_code < 300
|
||||
|
||||
def json(self) -> object:
|
||||
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 _RaiseOnSecondPostClient(_FakeAsyncClient):
|
||||
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)
|
||||
|
||||
|
||||
def test_create_github_pr_assigns_created_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(201, {"assignees": [{"login": "octocat"}]}),
|
||||
]
|
||||
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: add assignment",
|
||||
head_branch="feature",
|
||||
base_branch="main",
|
||||
body="body",
|
||||
assignee_login="octocat",
|
||||
)
|
||||
)
|
||||
|
||||
assert result == ("https://github.com/o/r/pull/12", 12, False)
|
||||
assert calls == [
|
||||
(
|
||||
"POST",
|
||||
"https://api.github.com/repos/o/r/pulls",
|
||||
{
|
||||
"title": "feat: add assignment",
|
||||
"head": "feature",
|
||||
"base": "main",
|
||||
"body": "body",
|
||||
"draft": True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"POST",
|
||||
"https://api.github.com/repos/o/r/issues/12/assignees",
|
||||
{"assignees": ["octocat"]},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def test_create_github_pr_assigns_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(201, {"assignees": [{"login": "octocat"}]}),
|
||||
]
|
||||
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: add assignment",
|
||||
head_branch="feature",
|
||||
base_branch="main",
|
||||
body="body",
|
||||
assignee_login="octocat",
|
||||
)
|
||||
)
|
||||
|
||||
assert result == ("https://github.com/o/r/pull/7", 7, True)
|
||||
assert calls == [
|
||||
(
|
||||
"POST",
|
||||
"https://api.github.com/repos/o/r/pulls",
|
||||
{
|
||||
"title": "feat: add assignment",
|
||||
"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},
|
||||
),
|
||||
(
|
||||
"POST",
|
||||
"https://api.github.com/repos/o/r/issues/7/assignees",
|
||||
{"assignees": ["octocat"]},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def test_create_github_pr_keeps_success_when_assignment_fails(
|
||||
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})]
|
||||
monkeypatch.setattr(
|
||||
github.httpx,
|
||||
"AsyncClient",
|
||||
lambda: _RaiseOnSecondPostClient(responses, calls),
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
github.create_github_pr(
|
||||
repo_owner="o",
|
||||
repo_name="r",
|
||||
github_token="token",
|
||||
title="feat: add assignment",
|
||||
head_branch="feature",
|
||||
base_branch="main",
|
||||
body="body",
|
||||
assignee_login="octocat",
|
||||
)
|
||||
)
|
||||
|
||||
assert result == ("https://github.com/o/r/pull/12", 12, False)
|
||||
assert calls == [
|
||||
(
|
||||
"POST",
|
||||
"https://api.github.com/repos/o/r/pulls",
|
||||
{
|
||||
"title": "feat: add assignment",
|
||||
"head": "feature",
|
||||
"base": "main",
|
||||
"body": "body",
|
||||
"draft": True,
|
||||
},
|
||||
),
|
||||
(
|
||||
"POST",
|
||||
"https://api.github.com/repos/o/r/issues/12/assignees",
|
||||
{"assignees": ["octocat"]},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def test_resolve_triggering_user_identity_keeps_github_login() -> None:
|
||||
identity = authorship.resolve_triggering_user_identity(
|
||||
{"configurable": {"github_login": "octocat", "github_user_id": 12345}}
|
||||
)
|
||||
|
||||
assert identity is not None
|
||||
assert identity.github_login == "octocat"
|
||||
assert identity.commit_email == "12345+octocat@users.noreply.github.com"
|
||||
Loading…
Add table
Reference in a new issue