From 150ff6f0c8ddef8c39c91a56185e6e2e9b3044c6 Mon Sep 17 00:00:00 2001 From: Aran Yogesh Date: Fri, 20 Mar 2026 10:53:33 -0700 Subject: [PATCH] fix: allow open-swe to be triggered on any GitHuh branch (#1100) * fix: allow open-swe to be triggered on any GitHuh branch * formmatting --- agent/middleware/open_pr.py | 15 +++++++++++++-- agent/server.py | 18 ++++++++++++++++++ agent/tools/commit_and_open_pr.py | 15 +++++++++++++-- agent/webapp.py | 25 +++++++++++++++++++++++-- 4 files changed, 67 insertions(+), 6 deletions(-) diff --git a/agent/middleware/open_pr.py b/agent/middleware/open_pr.py index 4e346387..116d6356 100644 --- a/agent/middleware/open_pr.py +++ b/agent/middleware/open_pr.py @@ -114,11 +114,22 @@ async def open_pr_if_needed( logger.info("Changes detected, preparing PR for thread %s", thread_id) + metadata = config.get("metadata", {}) + branch_name = metadata.get("branch_name") current_branch = await asyncio.to_thread(git_current_branch, sandbox_backend, repo_dir) - target_branch = f"open-swe/{thread_id}" + target_branch = branch_name if branch_name else f"open-swe/{thread_id}" if current_branch != target_branch: - await asyncio.to_thread(git_checkout_branch, sandbox_backend, repo_dir, target_branch) + if branch_name: + # Existing branch — plain checkout, do not create or reset + await asyncio.to_thread( + sandbox_backend.execute, + f"cd {repo_dir} && git checkout {target_branch}", + ) + else: + await asyncio.to_thread( + git_checkout_branch, sandbox_backend, repo_dir, target_branch + ) await asyncio.to_thread( git_config_user, diff --git a/agent/server.py b/agent/server.py index 5a0dd319..f7e49717 100644 --- a/agent/server.py +++ b/agent/server.py @@ -362,6 +362,24 @@ async def get_agent(config: RunnableConfig) -> Pregel: # noqa: PLR0915 msg = "Cannot proceed: no repo was cloned. Set 'repo.owner' and 'repo.name' in the configurable config" raise RuntimeError(msg) + branch_name = get_config().get("metadata", {}).get("branch_name") + if branch_name: + logger.info("Checking out branch '%s' in sandbox for thread %s", branch_name, thread_id) + loop = asyncio.get_event_loop() + safe_repo_dir = shlex.quote(repo_dir) + safe_branch = shlex.quote(branch_name) + checkout_result = await loop.run_in_executor( + None, + sandbox_backend.execute, + f"cd {safe_repo_dir} && git fetch origin && git checkout {safe_branch}", + ) + if checkout_result.exit_code != 0: + logger.warning( + "Failed to checkout branch '%s': %s", + branch_name, + checkout_result.output[:200] if checkout_result.output else "", + ) + linear_issue = config["configurable"].get("linear_issue", {}) linear_project_id = linear_issue.get("linear_project_id", "") linear_issue_number = linear_issue.get("linear_issue_number", "") diff --git a/agent/tools/commit_and_open_pr.py b/agent/tools/commit_and_open_pr.py index 87b7acd6..61755153 100644 --- a/agent/tools/commit_and_open_pr.py +++ b/agent/tools/commit_and_open_pr.py @@ -139,10 +139,21 @@ def commit_and_open_pr( if not (has_uncommitted_changes or has_unpushed_commits): return {"success": False, "error": "No changes detected", "pr_url": None} + metadata = config.get("metadata", {}) + branch_name = metadata.get("branch_name") current_branch = git_current_branch(sandbox_backend, repo_dir) - target_branch = f"open-swe/{thread_id}" + target_branch = branch_name if branch_name else f"open-swe/{thread_id}" if current_branch != target_branch: - if not git_checkout_branch(sandbox_backend, repo_dir, target_branch): + if branch_name: + # Existing branch — plain checkout, do not create or reset + result = sandbox_backend.execute(f"cd {repo_dir} && git checkout {target_branch}") + if result.exit_code != 0: + return { + "success": False, + "error": f"Failed to checkout branch {target_branch}", + "pr_url": None, + } + elif not git_checkout_branch(sandbox_backend, repo_dir, target_branch): return { "success": False, "error": f"Failed to checkout branch {target_branch}", diff --git a/agent/webapp.py b/agent/webapp.py index cf531640..e680117c 100644 --- a/agent/webapp.py +++ b/agent/webapp.py @@ -1253,8 +1253,29 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) -> thread_id = get_thread_id_from_branch(branch_name) if branch_name else None if not thread_id: - logger.warning("Could not extract thread_id from branch '%s', skipping", branch_name) - return + if not pr_number: + logger.warning( + "Could not determine thread_id for branch '%s' (no pr_number), skipping", + branch_name, + ) + return + owner = repo_config.get("owner", "") + name = repo_config.get("name", "") + stable_key = f"{owner}/{name}/pr/{pr_number}" + thread_id = str(uuid.uuid5(uuid.NAMESPACE_URL, stable_key)) + logger.info("Generated thread_id %s for non-open-swe branch '%s'", thread_id, branch_name) + langgraph_client = get_client(url=LANGGRAPH_URL) + try: + await langgraph_client.threads.update(thread_id, metadata={"branch_name": branch_name}) + except Exception as exc: # noqa: BLE001 + if _is_not_found_error(exc): + await langgraph_client.threads.create( + thread_id=thread_id, + if_exists="do_nothing", + metadata={"branch_name": branch_name}, + ) + else: + logger.warning("Failed to persist branch_name metadata for thread %s", thread_id) email = GITHUB_USER_EMAIL_MAP.get(github_login, "") if not email: