diff --git a/apps/agent/agent/middleware/open_pr.py b/apps/agent/agent/middleware/open_pr.py index 0dd29489..884aa570 100644 --- a/apps/agent/agent/middleware/open_pr.py +++ b/apps/agent/agent/middleware/open_pr.py @@ -18,15 +18,27 @@ from langgraph.config import get_config from langgraph.runtime import Runtime from ..encryption import decrypt_token -from ..utils.github import create_github_pr, get_github_default_branch +from ..utils.github import ( + create_github_pr, + get_github_default_branch, + git_add_all, + git_checkout_branch, + git_commit, + git_config_user, + git_current_branch, + git_fetch_origin, + git_has_uncommitted_changes, + git_has_unpushed_commits, + git_push, +) from ..utils.linear import comment_on_linear_issue from ..utils.sandbox_state import SANDBOX_BACKENDS logger = logging.getLogger(__name__) -def _extract_pr_params_from_messages(messages: list) -> dict[str, str] | None: - """Extract PR title/body/commit_message from the last commit_and_open_pr tool result.""" +def _extract_pr_params_from_messages(messages: list) -> dict[str, Any] | None: + """Extract commit_and_open_pr tool result payload.""" for msg in reversed(messages): if isinstance(msg, dict): content = msg.get("content", "") @@ -38,7 +50,7 @@ def _extract_pr_params_from_messages(messages: list) -> dict[str, str] | None: if name == "commit_and_open_pr" and content: try: parsed = _json.loads(content) if isinstance(content, str) else content - if isinstance(parsed, dict) and "title" in parsed: + if isinstance(parsed, dict): return parsed except (ValueError, TypeError): pass @@ -73,9 +85,9 @@ async def open_pr_if_needed( linear_issue = configurable.get("linear_issue", {}) linear_issue_id = linear_issue.get("id") - pr_params = _extract_pr_params_from_messages(messages) + pr_payload = _extract_pr_params_from_messages(messages) - if not pr_params: + if not pr_payload: logger.info("No commit_and_open_pr tool call found, skipping PR creation") if linear_issue_id and last_message_content: comment = f""" **Agent Response** @@ -84,9 +96,41 @@ async def open_pr_if_needed( await comment_on_linear_issue(linear_issue_id, comment) return None - pr_title = pr_params.get("title", "feat: Open SWE PR") - pr_body = pr_params.get("body", "Automated PR created by Open SWE agent.") - commit_message = pr_params.get("commit_message", pr_title) + if "success" in pr_payload: + pr_url = pr_payload.get("pr_url") + error = pr_payload.get("error") + if linear_issue_id and last_message_content: + if pr_url: + comment = f"""**Pull Request Created** + +I've created a pull request to address this issue: + +{pr_url} + +--- + **Agent Response** + +{last_message_content}""" + elif error: + comment = f"""**Pull Request Error** + +{error} + +--- + +**Agent Response** + +{last_message_content}""" + else: + comment = f""" **Agent Response** + +{last_message_content}""" + await comment_on_linear_issue(linear_issue_id, comment) + return None + + pr_title = pr_payload.get("title", "feat: Open SWE PR") + pr_body = pr_payload.get("body", "Automated PR created by Open SWE agent.") + commit_message = pr_payload.get("commit_message", pr_title) if not thread_id: if linear_issue_id and last_message_content: @@ -112,21 +156,14 @@ async def open_pr_if_needed( await comment_on_linear_issue(linear_issue_id, comment) return None - result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git status --porcelain" + has_uncommitted_changes = await asyncio.to_thread( + git_has_uncommitted_changes, sandbox_backend, repo_dir ) - has_uncommitted_changes = result.exit_code == 0 and result.output.strip() - - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git fetch origin 2>/dev/null || true" + await asyncio.to_thread(git_fetch_origin, sandbox_backend, repo_dir) + has_unpushed_commits = await asyncio.to_thread( + git_has_unpushed_commits, sandbox_backend, repo_dir ) - git_log_cmd = ( - f"cd {repo_dir} && git log --oneline @{{upstream}}..HEAD 2>/dev/null " - "|| git log --oneline origin/HEAD..HEAD 2>/dev/null || echo ''" - ) - unpushed_result = await asyncio.to_thread(sandbox_backend.execute, git_log_cmd) - has_unpushed_commits = unpushed_result.exit_code == 0 and unpushed_result.output.strip() has_changes = has_uncommitted_changes or has_unpushed_commits @@ -141,57 +178,36 @@ async def open_pr_if_needed( logger.info("Changes detected, preparing PR for thread %s", thread_id) - branch_result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git rev-parse --abbrev-ref HEAD" + current_branch = await asyncio.to_thread( + git_current_branch, sandbox_backend, repo_dir ) - current_branch = branch_result.output.strip() if branch_result.exit_code == 0 else "" target_branch = f"open-swe/{thread_id}" if current_branch != target_branch: - checkout_result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git checkout -b {target_branch}" + await asyncio.to_thread( + git_checkout_branch, sandbox_backend, repo_dir, target_branch ) - if checkout_result.exit_code != 0: - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git checkout {target_branch}" - ) await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git config user.name 'Open SWE[bot]'" - ) - await asyncio.to_thread( - sandbox_backend.execute, - f"cd {repo_dir} && git config user.email 'Open SWE@users.noreply.github.com'", - ) - - await asyncio.to_thread(sandbox_backend.execute, f"cd {repo_dir} && git add -A") - - safe_commit_msg = commit_message.replace("'", "'\\''") - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git commit -m '{safe_commit_msg}'" + git_config_user, + sandbox_backend, + repo_dir, + "Open SWE[bot]", + "Open SWE@users.noreply.github.com", ) + await asyncio.to_thread(git_add_all, sandbox_backend, repo_dir) + await asyncio.to_thread(git_commit, sandbox_backend, repo_dir, commit_message) encrypted_token = configurable.get("github_token_encrypted") + github_token = None if encrypted_token: github_token = decrypt_token(encrypted_token) if github_token: - remote_result = await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git remote get-url origin" + await asyncio.to_thread( + git_push, sandbox_backend, repo_dir, target_branch, github_token ) - if remote_result.exit_code == 0: - remote_url = remote_result.output.strip() - if "github.com" in remote_url and "@" not in remote_url: - auth_url = remote_url.replace("https://", f"https://git:{github_token}@") - await asyncio.to_thread( - sandbox_backend.execute, - f"cd {repo_dir} && git push {auth_url} {target_branch}", - ) - else: - await asyncio.to_thread( - sandbox_backend.execute, f"cd {repo_dir} && git push origin {target_branch}" - ) base_branch = await get_github_default_branch(repo_owner, repo_name, github_token) logger.info("Using base branch: %s", base_branch) diff --git a/apps/agent/agent/tools/commit_and_open_pr.py b/apps/agent/agent/tools/commit_and_open_pr.py index c29078c1..66dd3103 100644 --- a/apps/agent/agent/tools/commit_and_open_pr.py +++ b/apps/agent/agent/tools/commit_and_open_pr.py @@ -1,6 +1,26 @@ +import asyncio import logging from typing import Any +from langgraph.config import get_config + +from ..encryption import decrypt_token +from ..server import _create_langsmith_sandbox +from ..utils.github import ( + create_github_pr, + get_github_default_branch, + git_add_all, + git_checkout_branch, + git_commit, + git_config_user, + git_current_branch, + git_fetch_origin, + git_has_uncommitted_changes, + git_has_unpushed_commits, + git_push, +) +from ..utils.sandbox_state import SANDBOX_BACKENDS + logger = logging.getLogger(__name__) @@ -85,10 +105,111 @@ def commit_and_open_pr( commit_message: Optional git commit message. If not provided, the PR title is used. Returns: - Dictionary with the result of the operation including PR URL if successful. + Dictionary containing: + - 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 """ - return { - "title": title, - "body": body, - "commit_message": commit_message or title, - } + try: + config = get_config() + configurable = config.get("configurable", {}) + thread_id = configurable.get("thread_id") + if not thread_id: + return {"success": False, "error": "Missing thread_id in config", "pr_url": None} + + repo_config = configurable.get("repo", {}) + repo_owner = repo_config.get("owner") + repo_name = repo_config.get("name") + if not repo_owner or not repo_name: + return { + "success": False, + "error": "Missing repo owner/name in config", + "pr_url": None, + } + + sandbox_backend = SANDBOX_BACKENDS.get(thread_id) + if not sandbox_backend: + sandbox_id = configurable.get("sandbox_id") + + if not sandbox_id: + return {"success": False, "error": "No sandbox found for thread", "pr_url": None} + + sandbox_backend = _create_langsmith_sandbox(sandbox_id) + SANDBOX_BACKENDS[thread_id] = sandbox_backend + + repo_dir = f"/workspace/{repo_name}" + + has_uncommitted_changes = git_has_uncommitted_changes(sandbox_backend, repo_dir) + git_fetch_origin(sandbox_backend, repo_dir) + has_unpushed_commits = git_has_unpushed_commits(sandbox_backend, repo_dir) + + if not (has_uncommitted_changes or has_unpushed_commits): + return {"success": False, "error": "No changes detected", "pr_url": None} + + current_branch = git_current_branch(sandbox_backend, repo_dir) + target_branch = f"open-swe/{thread_id}" + if current_branch != target_branch: + if not git_checkout_branch(sandbox_backend, repo_dir, target_branch): + return { + "success": False, + "error": f"Failed to checkout branch {target_branch}", + "pr_url": None, + } + + git_config_user( + sandbox_backend, + repo_dir, + "Open SWE[bot]", + "Open SWE@users.noreply.github.com", + ) + git_add_all(sandbox_backend, repo_dir) + + commit_msg = commit_message or title + if has_uncommitted_changes: + commit_result = git_commit(sandbox_backend, repo_dir, commit_msg) + if commit_result.exit_code != 0: + return { + "success": False, + "error": f"Git commit failed: {commit_result.output.strip()}", + "pr_url": None, + } + + encrypted_token = configurable.get("github_token_encrypted") + github_token = decrypt_token(encrypted_token) if encrypted_token else None + if not github_token: + return {"success": False, "error": "Missing GitHub token", "pr_url": None} + + push_result = git_push(sandbox_backend, repo_dir, target_branch, github_token) + if push_result.exit_code != 0: + return { + "success": False, + "error": f"Git push failed: {push_result.output.strip()}", + "pr_url": None, + } + + base_branch = asyncio.run( + get_github_default_branch(repo_owner, repo_name, github_token) + ) + pr_url, _pr_number = asyncio.run( + create_github_pr( + repo_owner=repo_owner, + repo_name=repo_name, + github_token=github_token, + title=title, + head_branch=target_branch, + base_branch=base_branch, + body=body, + ) + ) + + if not pr_url: + return { + "success": False, + "error": "Failed to create GitHub PR", + "pr_url": None, + } + + return {"success": True, "error": None, "pr_url": pr_url} + 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 216cf581..c3e1e3cb 100644 --- a/apps/agent/agent/utils/github.py +++ b/apps/agent/agent/utils/github.py @@ -1,8 +1,10 @@ -"""GitHub API utilities.""" +"""GitHub API and git utilities.""" from __future__ import annotations import logging +import shlex +from typing import Any import httpx @@ -13,6 +15,97 @@ HTTP_CREATED = 201 HTTP_UNPROCESSABLE_ENTITY = 422 +def _run_git(sandbox_backend: Any, repo_dir: str, command: str) -> Any: + """Run a git command in the sandbox repo directory.""" + return sandbox_backend.execute(f"cd {repo_dir} && {command}") + + +def git_has_uncommitted_changes(sandbox_backend: Any, repo_dir: str) -> bool: + """Check whether the repo has uncommitted changes.""" + result = _run_git(sandbox_backend, repo_dir, "git status --porcelain") + return result.exit_code == 0 and bool(result.output.strip()) + + +def git_fetch_origin(sandbox_backend: Any, repo_dir: str) -> Any: + """Fetch latest from origin (best-effort).""" + return _run_git(sandbox_backend, repo_dir, "git fetch origin 2>/dev/null || true") + + +def git_has_unpushed_commits(sandbox_backend: Any, repo_dir: str) -> bool: + """Check whether there are commits not pushed to upstream.""" + git_log_cmd = ( + "git log --oneline @{upstream}..HEAD 2>/dev/null " + "|| git log --oneline origin/HEAD..HEAD 2>/dev/null || echo ''" + ) + result = _run_git(sandbox_backend, repo_dir, git_log_cmd) + return result.exit_code == 0 and bool(result.output.strip()) + + +def git_current_branch(sandbox_backend: Any, repo_dir: str) -> str: + """Get the current git branch name.""" + result = _run_git(sandbox_backend, repo_dir, "git rev-parse --abbrev-ref HEAD") + return result.output.strip() if result.exit_code == 0 else "" + + +def git_checkout_branch(sandbox_backend: Any, repo_dir: str, branch: str) -> bool: + """Checkout branch, creating it if needed.""" + safe_branch = shlex.quote(branch) + checkout_result = _run_git( + sandbox_backend, repo_dir, f"git checkout -b {safe_branch}" + ) + if checkout_result.exit_code == 0: + return True + fallback = _run_git(sandbox_backend, repo_dir, f"git checkout {safe_branch}") + return fallback.exit_code == 0 + + +def git_config_user( + sandbox_backend: Any, + repo_dir: str, + name: str, + email: str, +) -> None: + """Configure git user name and email.""" + safe_name = shlex.quote(name) + safe_email = shlex.quote(email) + _run_git(sandbox_backend, repo_dir, f"git config user.name {safe_name}") + _run_git(sandbox_backend, repo_dir, f"git config user.email {safe_email}") + + +def git_add_all(sandbox_backend: Any, repo_dir: str) -> Any: + """Stage all changes.""" + return _run_git(sandbox_backend, repo_dir, "git add -A") + + +def git_commit(sandbox_backend: Any, repo_dir: str, message: str) -> Any: + """Commit staged changes with the given message.""" + safe_message = shlex.quote(message) + return _run_git(sandbox_backend, repo_dir, f"git commit -m {safe_message}") + + +def git_get_remote_url(sandbox_backend: Any, repo_dir: str) -> str | None: + """Get the origin remote URL.""" + result = _run_git(sandbox_backend, repo_dir, "git remote get-url origin") + if result.exit_code != 0: + return None + return result.output.strip() + + +def git_push( + sandbox_backend: Any, + repo_dir: str, + branch: str, + github_token: str | None = None, +) -> Any: + """Push the branch to origin, using a token if needed.""" + safe_branch = shlex.quote(branch) + remote_url = git_get_remote_url(sandbox_backend, repo_dir) + if remote_url and "github.com" in remote_url and "@" not in remote_url and github_token: + auth_url = remote_url.replace("https://", f"https://git:{github_token}@") + return _run_git(sandbox_backend, repo_dir, f"git push {auth_url} {safe_branch}") + return _run_git(sandbox_backend, repo_dir, f"git push origin {safe_branch}") + + async def create_github_pr( repo_owner: str, repo_name: str,