diff --git a/apps/agent/agent/middleware/open_pr.py b/apps/agent/agent/middleware/open_pr.py index 884aa570..158b7671 100644 --- a/apps/agent/agent/middleware/open_pr.py +++ b/apps/agent/agent/middleware/open_pr.py @@ -98,12 +98,15 @@ async def open_pr_if_needed( if "success" in pr_payload: pr_url = pr_payload.get("pr_url") + pr_existing = bool(pr_payload.get("pr_existing", False)) error = pr_payload.get("error") if linear_issue_id and last_message_content: if pr_url: - comment = f"""**Pull Request Created** + header = "Pull Request Updated" if pr_existing else "Pull Request Created" + action = "updated the existing" if pr_existing else "created a" + comment = f"""**{header}** -I've created a pull request to address this issue: +I've {action} pull request to address this issue: {pr_url} @@ -212,7 +215,7 @@ I've created a pull request to address this issue: base_branch = await get_github_default_branch(repo_owner, repo_name, github_token) logger.info("Using base branch: %s", base_branch) - pr_url, pr_number = await create_github_pr( + pr_url, pr_number, pr_existing = await create_github_pr( repo_owner=repo_owner, repo_name=repo_name, github_token=github_token, @@ -227,9 +230,11 @@ I've created a pull request to address this issue: if linear_issue_id and last_message_content: if pr_url: - comment = f"""**Pull Request Created** + header = "Pull Request Updated" if pr_existing else "Pull Request Created" + action = "updated the existing" if pr_existing else "created a" + comment = f"""**{header}** -I've created a pull request to address this issue: +I've {action} pull request to address this issue: **[PR #{pr_number}: {pr_title}]({pr_url})** diff --git a/apps/agent/agent/tools/commit_and_open_pr.py b/apps/agent/agent/tools/commit_and_open_pr.py index c6c9fe92..6a40ac73 100644 --- a/apps/agent/agent/tools/commit_and_open_pr.py +++ b/apps/agent/agent/tools/commit_and_open_pr.py @@ -109,6 +109,7 @@ def commit_and_open_pr( - success: Whether the operation completed successfully - error: Error string if something failed, otherwise None - pr_url: URL of the created PR if successful, otherwise None + - pr_existing: Whether a PR already existed for this branch """ try: config = get_config() @@ -190,7 +191,7 @@ def commit_and_open_pr( base_branch = asyncio.run( get_github_default_branch(repo_owner, repo_name, github_token) ) - pr_url, _pr_number = asyncio.run( + pr_url, _pr_number, pr_existing = asyncio.run( create_github_pr( repo_owner=repo_owner, repo_name=repo_name, @@ -207,9 +208,15 @@ def commit_and_open_pr( "success": False, "error": "Failed to create GitHub PR", "pr_url": None, + "pr_existing": False, } - return {"success": True, "error": None, "pr_url": pr_url} + return { + "success": True, + "error": None, + "pr_url": pr_url, + "pr_existing": pr_existing, + } except Exception as e: logger.exception("commit_and_open_pr failed") return {"success": False, "error": f"{type(e).__name__}: {e}", "pr_url": None} diff --git a/apps/agent/agent/utils/github.py b/apps/agent/agent/utils/github.py index c3e1e3cb..bee5ee66 100644 --- a/apps/agent/agent/utils/github.py +++ b/apps/agent/agent/utils/github.py @@ -114,7 +114,7 @@ async def create_github_pr( head_branch: str, base_branch: str, body: str, -) -> tuple[str | None, int | None]: +) -> tuple[str | None, int | None, bool]: """Create a GitHub pull request via the API. Args: @@ -127,7 +127,7 @@ async def create_github_pr( body: PR description Returns: - Tuple of (pr_url, pr_number) if successful, (None, None) otherwise + Tuple of (pr_url, pr_number, pr_existing) if successful, (None, None, False) otherwise """ pr_payload = { "title": title, @@ -144,8 +144,8 @@ async def create_github_pr( repo_name, ) - try: - async with httpx.AsyncClient() as http_client: + 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={ @@ -162,10 +162,20 @@ async def create_github_pr( 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 + 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", @@ -176,11 +186,41 @@ async def create_github_pr( if "errors" in pr_data: logger.error("GitHub API errors detail: %s", pr_data.get("errors")) - return None, None + return None, None, False - except httpx.HTTPError: - logger.exception("Failed to create PR via GitHub API") - return None, None + except httpx.HTTPError: + logger.exception("Failed to create PR via GitHub API") + return None, None, False + + +async def _find_existing_pr( + http_client: httpx.AsyncClient, + repo_owner: str, + repo_name: str, + github_token: str, + head_branch: str, +) -> tuple[str | None, int | None]: + """Find an existing PR for the given head branch.""" + headers = { + "Authorization": f"Bearer {github_token}", + "Accept": "application/vnd.github+json", + "X-GitHub-Api-Version": "2022-11-28", + } + head_ref = f"{repo_owner}:{head_branch}" + for state in ("open", "all"): + response = await http_client.get( + f"https://api.github.com/repos/{repo_owner}/{repo_name}/pulls", + headers=headers, + params={"head": head_ref, "state": state, "per_page": 1}, + ) + if response.status_code != 200: # noqa: PLR2004 + continue + data = response.json() + if not data: + continue + pr = data[0] + return pr.get("html_url"), pr.get("number") + return None, None async def get_github_default_branch( diff --git a/apps/agent/agent/webapp.py b/apps/agent/agent/webapp.py index 917cb9e4..04c01590 100644 --- a/apps/agent/agent/webapp.py +++ b/apps/agent/agent/webapp.py @@ -575,6 +575,9 @@ async def process_linear_issue( # noqa: PLR0912, PLR0915 bot_message_prefixes = ( "🔐 **GitHub Authentication Required**", "✅ **Pull Request Created**", + "✅ **Pull Request Updated**", + "**Pull Request Created**", + "**Pull Request Updated**", "🤖 **Agent Response**", "❌ **Agent Error**", ) @@ -759,6 +762,9 @@ async def linear_webhook( # noqa: PLR0911, PLR0912, PLR0915 bot_message_prefixes = [ "🔐 **GitHub Authentication Required**", "✅ **Pull Request Created**", + "✅ **Pull Request Updated**", + "**Pull Request Created**", + "**Pull Request Updated**", "🤖 **Agent Response**", "❌ **Agent Error**", ]