diff --git a/agent/middleware/open_pr.py b/agent/middleware/open_pr.py index 12d9325e..27f5dccf 100644 --- a/agent/middleware/open_pr.py +++ b/agent/middleware/open_pr.py @@ -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") diff --git a/agent/tools/commit_and_open_pr.py b/agent/tools/commit_and_open_pr.py index b4922c91..da17650c 100644 --- a/agent/tools/commit_and_open_pr.py +++ b/agent/tools/commit_and_open_pr.py @@ -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, ) ) diff --git a/agent/utils/authorship.py b/agent/utils/authorship.py index 224a0bcf..5bd0bc76 100644 --- a/agent/utils/authorship.py +++ b/agent/utils/authorship.py @@ -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", {}) diff --git a/agent/utils/github.py b/agent/utils/github.py index be7fc835..5c81d24f 100644 --- a/agent/utils/github.py +++ b/agent/utils/github.py @@ -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, diff --git a/tests/test_github_pr_assignment.py b/tests/test_github_pr_assignment.py new file mode 100644 index 00000000..693d1b13 --- /dev/null +++ b/tests/test_github_pr_assignment.py @@ -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"