fix: Auto assign PRs to creator (#1211)

* fix: Auto assign PRs to creator

* cr
This commit is contained in:
Brace Sproul 2026-04-21 14:13:55 -07:00 • committed by GitHub
parent 5925a90a95
commit e9b94ac8ba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 282 additions and 0 deletions

View file

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

View file

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

View file

@ -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", {})

View file

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

View 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"