mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 15:03:16 +00:00
Merge pull request #935 from langchain-ai/yogesh/centeralize-gitops
refactor: centralize git ops and implement commit_and_open_pr flow
This commit is contained in:
commit
fdee31eb8b
3 changed files with 293 additions and 63 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue