fix: allow open-swe to be triggered on any GitHuh branch (#1100)

* fix: allow open-swe to be triggered on any GitHuh branch

* formmatting
This commit is contained in:
Aran Yogesh 2026-03-20 10:53:33 -07:00 • committed by GitHub
parent ea27978d51
commit 150ff6f0c8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 67 additions and 6 deletions

View file

@ -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,

View file

@ -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", "")

View file

@ -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}",

View file

@ -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: