feat: recover thread work as patch (#1615)

* feat: recover thread work as patch

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

* fix: search sandbox cwd for recovery patches

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>

---------

Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
Johannes du Plessis 2026-06-25 17:18:16 -07:00 • committed by GitHub
parent 29015fadc4
commit f29868ffe7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 495 additions and 11 deletions

View file

@ -155,6 +155,7 @@ from .thread_api import (
delete_dashboard_thread,
get_dashboard_thread,
get_dashboard_thread_pr_diff,
get_dashboard_thread_recovery_patch,
get_dashboard_thread_state,
list_dashboard_threads,
list_dashboard_threads_page,
@ -1534,6 +1535,23 @@ async def api_get_thread(
)
@router.get("/threads/{thread_id}/recovery.patch")
async def api_get_thread_recovery_patch(
thread_id: str,
session: dict[str, Any] = _SESSION_DEP,
) -> Response:
content, filename = await get_dashboard_thread_recovery_patch(
thread_id,
session["sub"],
email=session.get("email"),
)
return Response(
content=content,
media_type="text/x-diff",
headers={"Content-Disposition": f'attachment; filename="{filename}"'},
)
@router.get("/threads/{thread_id}/pr-diff")
async def api_get_thread_pr_diff(
thread_id: str,

View file

@ -18,6 +18,7 @@ from langchain_core.messages.content import create_image_block
from pydantic import BaseModel, ConfigDict, Field
from ..utils.langsmith import get_langsmith_trace_url
from ..utils.sandbox import create_sandbox
from ..utils.thread_ops import (
get_thread_active_status,
langgraph_client,
@ -56,6 +57,8 @@ _PROXY_STREAM_TIMEOUT = httpx.Timeout(None)
_SURFACED_SOURCES: tuple[str, ...] = ("dashboard", "github", "slack", "linear", "schedule")
# PR lifecycle states surfaced to the UI for a thread's associated pull request.
_PR_STATES: frozenset[str] = frozenset({"draft", "open", "merged", "closed"})
_RECOVERY_PATCH_LIMIT_BYTES = 25 * 1024 * 1024
_RECOVERY_PATCH_TIMEOUT_SECONDS = 120
def _agent_version_metadata() -> dict[str, str]:
@ -1382,6 +1385,247 @@ async def get_dashboard_thread_state(
return result
def _recovery_patch_filename(thread_id: str) -> str:
safe = "".join(c if c.isalnum() or c in {"-", "_", "."} else "-" for c in thread_id)
return f"open-swe-{(safe or 'thread')[:80]}.patch"
def _response_output(result: Any) -> str:
output = result.get("output") if isinstance(result, dict) else getattr(result, "output", "")
return output if isinstance(output, str) else str(output or "")
def _response_exit_code(result: Any) -> int | None:
value = (
result.get("exit_code") if isinstance(result, dict) else getattr(result, "exit_code", None)
)
return value if isinstance(value, int) else None
def _download_content(result: Any) -> bytes | None:
for attr in ("content", "data", "bytes"):
value = result.get(attr) if isinstance(result, dict) else getattr(result, attr, None)
if isinstance(value, bytes):
return value
if isinstance(value, str):
return value.encode()
file_data = (
result.get("file_data") if isinstance(result, dict) else getattr(result, "file_data", None)
)
if isinstance(file_data, bytes):
return file_data
if isinstance(file_data, str):
return file_data.encode()
if isinstance(file_data, dict):
for key in ("content", "data", "bytes"):
value = file_data.get(key)
if isinstance(value, bytes):
return value
if isinstance(value, str):
return value.encode()
return None
def _recovery_patch_command(metadata: dict[str, Any], thread_id: str) -> str:
_, name, _ = _metadata_repo(metadata)
payload = {
"repo_name": name,
"base_branch": metadata.get("base_branch")
if isinstance(metadata.get("base_branch"), str)
else "main",
"thread_key": _recovery_patch_filename(thread_id).removesuffix(".patch"),
}
encoded = base64.b64encode(json.dumps(payload).encode()).decode()
script = r"""python - <<'PY'
import base64
import json
import subprocess
import sys
from pathlib import Path
PAYLOAD = json.loads(base64.b64decode('__PAYLOAD__').decode())
WORKSPACE_FALLBACK = Path('/workspace')
def git(repo, args, check=True):
result = subprocess.run(
['git', '-C', str(repo), *args],
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
if check and result.returncode != 0:
detail = result.stderr.decode(errors='replace').strip()
raise RuntimeError(detail or 'git ' + ' '.join(args) + ' failed')
return result
def search_roots():
roots = [Path.cwd().resolve(), WORKSPACE_FALLBACK]
seen = set()
for root in roots:
if root in seen:
continue
seen.add(root)
if root.exists():
yield root
def repo_paths():
repo_name = PAYLOAD.get('repo_name')
for root in search_roots():
if isinstance(repo_name, str) and repo_name:
yield root / Path(repo_name).name
yield root
for child in sorted(root.iterdir()):
if child.is_dir():
yield child
def find_repo():
seen = set()
for path in repo_paths():
if path in seen:
continue
seen.add(path)
if not (path / '.git').exists():
continue
result = git(path, ['rev-parse', '--show-toplevel'], check=False)
if result.returncode == 0:
root = Path(result.stdout.decode(errors='replace').strip())
if root.exists():
return root
raise RuntimeError('no git repository found in sandbox workspace')
def safe_ref(value):
if not isinstance(value, str) or not value or len(value) > 200:
return None
if value.startswith('-') or '\x00' in value or '\n' in value or '\r' in value:
return None
return value
def commit_for(repo, ref):
result = git(repo, ['rev-parse', '--verify', ref + '^{commit}'], check=False)
if result.returncode == 0:
return result.stdout.decode(errors='replace').strip()
return None
def merge_base(repo):
base_branch = safe_ref(PAYLOAD.get('base_branch')) or 'main'
refs = ['origin/' + base_branch, base_branch, 'origin/main', 'main', 'origin/master', 'master', 'HEAD~1']
for ref in refs:
commit = commit_for(repo, ref)
if not commit:
continue
result = git(repo, ['merge-base', 'HEAD', commit], check=False)
if result.returncode == 0:
return result.stdout.decode(errors='replace').strip()
return commit
return git(repo, ['hash-object', '-t', 'tree', '/dev/null']).stdout.decode(errors='replace').strip()
def write_patch(repo, base):
patch_path = Path('/tmp') / ((PAYLOAD.get('thread_key') or 'open-swe-recovery') + '.patch')
with patch_path.open('wb') as patch_file:
tracked = git(repo, ['diff', '--binary', '--full-index', base, '--', '.']).stdout
patch_file.write(tracked)
untracked = git(repo, ['ls-files', '--others', '--exclude-standard', '-z']).stdout
for raw_path in [p for p in untracked.split(b'\0') if p]:
rel_path = raw_path.decode('utf-8', errors='surrogateescape')
full_path = repo / rel_path
if not full_path.is_file():
continue
result = git(
repo,
['diff', '--no-index', '--binary', '--full-index', '--', '/dev/null', rel_path],
check=False,
)
if result.returncode not in {0, 1}:
detail = result.stderr.decode(errors='replace').strip()
raise RuntimeError(detail or 'failed to diff untracked file ' + rel_path)
if result.stdout:
if patch_file.tell() and not result.stdout.startswith(b'\n'):
patch_file.write(b'\n')
patch_file.write(result.stdout)
return patch_path
try:
repo = find_repo()
base = merge_base(repo)
patch_path = write_patch(repo, base)
print(json.dumps({'ok': True, 'path': str(patch_path), 'size': patch_path.stat().st_size}))
except Exception as exc:
print(json.dumps({'ok': False, 'error': str(exc)}))
sys.exit(1)
PY"""
return script.replace("__PAYLOAD__", encoded)
async def get_dashboard_thread_recovery_patch(
thread_id: str, login: str, *, email: str | None = None
) -> tuple[bytes, str]:
thread = await _authorized_thread(thread_id, login, email=email)
metadata = thread.get("metadata") if isinstance(thread.get("metadata"), dict) else {}
sandbox_id = metadata.get("sandbox_id")
if not isinstance(sandbox_id, str) or not sandbox_id:
raise HTTPException(404, "thread has no recoverable sandbox")
try:
sandbox = await asyncio.to_thread(create_sandbox, sandbox_id)
except Exception as exc: # noqa: BLE001
logger.debug("Could not connect to sandbox %s for recovery", sandbox_id, exc_info=True)
raise HTTPException(502, "could not connect to thread sandbox") from exc
try:
result = await asyncio.to_thread(
sandbox.execute,
_recovery_patch_command(metadata, thread_id),
timeout=_RECOVERY_PATCH_TIMEOUT_SECONDS,
)
except Exception as exc: # noqa: BLE001
logger.debug("Recovery patch generation failed for %s", thread_id, exc_info=True)
raise HTTPException(502, "failed to generate recovery patch") from exc
output = _response_output(result).strip()
try:
payload = json.loads(output.splitlines()[-1])
except (IndexError, json.JSONDecodeError) as exc:
logger.debug("Invalid recovery patch response for %s: %s", thread_id, output)
raise HTTPException(502, "failed to generate recovery patch") from exc
if _response_exit_code(result) not in {0, None} or payload.get("ok") is not True:
detail = payload.get("error") if isinstance(payload.get("error"), str) else None
logger.debug("Recovery patch generation failed for %s: %s", thread_id, detail)
raise HTTPException(502, detail or "failed to generate recovery patch")
size = payload.get("size")
if not isinstance(size, int):
raise HTTPException(502, "failed to generate recovery patch")
if size == 0:
raise HTTPException(404, "thread has no recoverable changes")
if size > _RECOVERY_PATCH_LIMIT_BYTES:
raise HTTPException(413, "recovery patch is too large to download")
patch_path = payload.get("path")
if not isinstance(patch_path, str) or not patch_path.startswith("/tmp/"):
raise HTTPException(502, "failed to generate recovery patch")
try:
downloads = await asyncio.to_thread(sandbox.download_files, [patch_path])
except Exception as exc: # noqa: BLE001
logger.debug("Recovery patch download failed for %s", thread_id, exc_info=True)
raise HTTPException(502, "failed to download recovery patch") from exc
if not downloads:
raise HTTPException(502, "failed to download recovery patch")
content = _download_content(downloads[0])
if content is None:
raise HTTPException(502, "failed to download recovery patch")
return content, _recovery_patch_filename(thread_id)
# No app-installation-token fallback: PR file contents must be fetched with
# the user's own credential so GitHub enforces their current repo access.
async def _github_token_for_login(login: str) -> str:

View file

@ -1,4 +1,6 @@
import base64
import json
from types import SimpleNamespace
import pytest
from fastapi import HTTPException
@ -304,6 +306,132 @@ def test_thread_summary_omits_pr_when_no_pr_metadata() -> None:
assert "diffStats" not in summary
async def test_recovery_patch_requires_thread_owner(monkeypatch) -> None:
class FakeThreads:
async def get(self, thread_id: str) -> dict[str, object]:
return {
"thread_id": thread_id,
"metadata": {"source": "dashboard", "github_login": "owner", "sandbox_id": "sbx"},
}
class FakeClient:
threads = FakeThreads()
monkeypatch.setattr(thread_api, "langgraph_client", lambda: FakeClient())
with pytest.raises(HTTPException) as exc_info:
await thread_api.get_dashboard_thread_recovery_patch("tid", "intruder")
assert exc_info.value.status_code == 404
async def test_recovery_patch_requires_sandbox(monkeypatch) -> None:
async def fake_authorized_thread(thread_id: str, login: str, *, email: str | None = None):
return {"thread_id": thread_id, "metadata": {"source": "dashboard", "github_login": login}}
monkeypatch.setattr(thread_api, "_authorized_thread", fake_authorized_thread)
with pytest.raises(HTTPException) as exc_info:
await thread_api.get_dashboard_thread_recovery_patch("tid", "octocat")
assert exc_info.value.status_code == 404
assert "sandbox" in exc_info.value.detail
async def test_recovery_patch_downloads_generated_patch(monkeypatch) -> None:
async def fake_authorized_thread(thread_id: str, login: str, *, email: str | None = None):
return {
"thread_id": thread_id,
"metadata": {
"source": "dashboard",
"github_login": login,
"sandbox_id": "sbx",
"repo_owner": "octo",
"repo_name": "repo",
"base_branch": "main",
},
}
class FakeSandbox:
def execute(self, command: str, *, timeout: int | None = None):
assert "repo" in command
assert timeout == thread_api._RECOVERY_PATCH_TIMEOUT_SECONDS
return SimpleNamespace(
output=json.dumps({"ok": True, "path": "/tmp/open-swe-tid.patch", "size": 11}),
exit_code=0,
)
def download_files(self, paths: list[str]):
assert paths == ["/tmp/open-swe-tid.patch"]
return [SimpleNamespace(content=b"patch bytes")]
monkeypatch.setattr(thread_api, "_authorized_thread", fake_authorized_thread)
monkeypatch.setattr(thread_api, "create_sandbox", lambda sandbox_id: FakeSandbox())
content, filename = await thread_api.get_dashboard_thread_recovery_patch("tid", "octocat")
assert content == b"patch bytes"
assert filename == "open-swe-tid.patch"
async def test_recovery_patch_rejects_empty_patch(monkeypatch) -> None:
async def fake_authorized_thread(thread_id: str, login: str, *, email: str | None = None):
return {"thread_id": thread_id, "metadata": {"sandbox_id": "sbx", "github_login": login}}
class FakeSandbox:
def execute(self, command: str, *, timeout: int | None = None):
return SimpleNamespace(
output=json.dumps({"ok": True, "path": "/tmp/open-swe-tid.patch", "size": 0}),
exit_code=0,
)
monkeypatch.setattr(thread_api, "_authorized_thread", fake_authorized_thread)
monkeypatch.setattr(thread_api, "create_sandbox", lambda sandbox_id: FakeSandbox())
with pytest.raises(HTTPException) as exc_info:
await thread_api.get_dashboard_thread_recovery_patch("tid", "octocat")
assert exc_info.value.status_code == 404
assert "changes" in exc_info.value.detail
async def test_recovery_patch_enforces_size_limit(monkeypatch) -> None:
async def fake_authorized_thread(thread_id: str, login: str, *, email: str | None = None):
return {"thread_id": thread_id, "metadata": {"sandbox_id": "sbx", "github_login": login}}
class FakeSandbox:
def execute(self, command: str, *, timeout: int | None = None):
return SimpleNamespace(
output=json.dumps(
{
"ok": True,
"path": "/tmp/open-swe-tid.patch",
"size": thread_api._RECOVERY_PATCH_LIMIT_BYTES + 1,
}
),
exit_code=0,
)
monkeypatch.setattr(thread_api, "_authorized_thread", fake_authorized_thread)
monkeypatch.setattr(thread_api, "create_sandbox", lambda sandbox_id: FakeSandbox())
with pytest.raises(HTTPException) as exc_info:
await thread_api.get_dashboard_thread_recovery_patch("tid", "octocat")
assert exc_info.value.status_code == 413
def test_recovery_patch_searches_command_cwd_before_workspace_fallback() -> None:
command = thread_api._recovery_patch_command(
{"repo_name": "repo", "base_branch": "main"},
"tid",
)
assert "Path.cwd().resolve()" in command
assert "WORKSPACE_FALLBACK = Path('/workspace')" in command
assert "roots = [Path.cwd().resolve(), WORKSPACE_FALLBACK]" in command
async def test_proxy_commands_lazily_creates_missing_thread_only_for_run_start(
monkeypatch,
) -> None:

View file

@ -23,6 +23,7 @@ import type { GitStatus, GitStatusEntry } from "@pierre/trees"
import type { AgentThread, Message } from "@/lib/agents/types"
import type { ThreadPrDiffFile } from "@/lib/agents/api"
import type { ChangedFileSummaryItem } from "@/components/agents/messages"
import { agentsApi } from "@/lib/agents/api"
import { useAgentThreadPrDiff } from "@/lib/agents/queries"
import { ReviewTab } from "@/components/agents/ReviewTab"
import { buttonVariants } from "@/components/ui/button"
@ -346,6 +347,34 @@ export function AgentGitPanel({
}
const prDiff = useAgentThreadPrDiff(thread.id, Boolean(pr))
const [recoveringPatch, setRecoveringPatch] = useState(false)
const [recoveryError, setRecoveryError] = useState<string | null>(null)
const canDownloadRecovery =
thread.status !== "running" && thread.isOwner !== false
const downloadRecoveryPatch = useCallback(async () => {
setRecoveringPatch(true)
setRecoveryError(null)
try {
const { blob, filename } = await agentsApi.downloadThreadRecoveryPatch(
thread.id
)
const url = window.URL.createObjectURL(blob)
const link = document.createElement("a")
link.href = url
link.download = filename
document.body.appendChild(link)
link.click()
link.remove()
window.URL.revokeObjectURL(url)
} catch (error) {
setRecoveryError(
error instanceof Error ? error.message : "Failed to download patch"
)
} finally {
setRecoveringPatch(false)
}
}, [thread.id])
const chunks = useMemo(
() => messages.flatMap((message) => message.chunks),
@ -545,19 +574,42 @@ export function AgentGitPanel({
{label}
</button>
))}
{files.length > 0 && (
<span className="ml-auto flex items-center gap-2 text-[11px] text-[var(--ui-text-dim)]">
<span>
{files.length} file{files.length === 1 ? "" : "s"}
<div className="ml-auto flex min-w-0 items-center gap-2">
{recoveryError && (
<span
title={recoveryError}
className="max-w-40 truncate text-[11px] text-[var(--ui-danger)]"
>
{recoveryError}
</span>
<span className="text-[var(--ui-success)]">
+{totals.additions}
)}
{canDownloadRecovery && (
<button
type="button"
onClick={downloadRecoveryPatch}
disabled={recoveringPatch}
className={cn(
buttonVariants({ variant: "outline", size: "sm" }),
"h-7 px-2 text-[11px]"
)}
>
{recoveringPatch ? "Preparing…" : "Download patch"}
</button>
)}
{files.length > 0 && (
<span className="flex items-center gap-2 text-[11px] text-[var(--ui-text-dim)]">
<span>
{files.length} file{files.length === 1 ? "" : "s"}
</span>
<span className="text-[var(--ui-success)]">
+{totals.additions}
</span>
<span className="text-[var(--ui-danger)]">
-{totals.deletions}
</span>
</span>
<span className="text-[var(--ui-danger)]">
-{totals.deletions}
</span>
</span>
)}
)}
</div>
</div>
<div className="flex min-h-0 flex-1">

View file

@ -58,6 +58,11 @@ export interface ThreadPrDiff {
files: Array<ThreadPrDiffFile>
}
export interface ThreadRecoveryPatch {
blob: Blob
filename: string
}
export interface ThreadsPageParams {
limit?: number
offset?: number
@ -125,6 +130,39 @@ async function agentsRequest<T>(
return (await res.json()) as T
}
function filenameFromContentDisposition(value: string | null): string | null {
const match = /filename="([^"]+)"/.exec(value ?? "")
return match?.[1] ?? null
}
async function agentsBlobRequest(path: string): Promise<ThreadRecoveryPatch> {
const res = await fetch(`${API_BASE}/dashboard/api${path}`, {
credentials: "include",
headers: { Accept: "text/x-diff" },
})
if (!res.ok) {
let message = res.statusText
try {
const body = await res.json()
if (body?.detail) {
message =
typeof body.detail === "string"
? body.detail
: JSON.stringify(body.detail)
}
} catch {
/* ignore */
}
throw new AgentsApiError(res.status, message)
}
return {
blob: await res.blob(),
filename:
filenameFromContentDisposition(res.headers.get("content-disposition")) ??
"open-swe-recovery.patch",
}
}
function buildThreadsPageQuery(params: ThreadsPageParams): string {
const search = new URLSearchParams()
if (params.limit != null) search.set("limit", String(params.limit))
@ -219,6 +257,10 @@ export const agentsApi = {
agentsRequest<ThreadPrDiff>(
`/threads/${encodeURIComponent(threadId)}/pr-diff`
),
downloadThreadRecoveryPatch: (threadId: string) =>
agentsBlobRequest(
`/threads/${encodeURIComponent(threadId)}/recovery.patch`
),
streamUrl: (threadId: string) =>
`${API_BASE}/dashboard/api/threads/${encodeURIComponent(threadId)}/stream`,
}