diff --git a/agent/middleware/workflow_push_guard.py b/agent/middleware/workflow_push_guard.py index f8afce39..6d4bd284 100644 --- a/agent/middleware/workflow_push_guard.py +++ b/agent/middleware/workflow_push_guard.py @@ -300,21 +300,10 @@ def _run_coroutine_sync(coro: Awaitable[ToolMessage | Command]) -> ToolMessage | return value -def _workflow_tree_at_head( - backend: Any, repo_dir: str | None, head: str -) -> tuple[list[str], str] | None: - """Return the sorted workflow file paths and a stable hash of the head workflow tree. - - Uses `git ls-tree -r` so the hash binds to the blob SHAs at the pushed head, not to - a diff against a sandbox-writable base ref. This prevents a confused deputy where the - sandbox rewrites `refs/remotes/origin/*` to hide a malicious workflow file from the - approved diff while still deploying it. - """ - ls_tree = _run_git(backend, repo_dir, f"ls-tree -r {shlex.quote(head)} -- .github/workflows") - if not ls_tree.ok: - return None +def _parse_ls_tree(output: str) -> list[tuple[str, str]]: + """Parse `git ls-tree -r` output into a sorted list of (sha, path) tuples.""" entries: list[tuple[str, str]] = [] - for line in ls_tree.output.splitlines(): + for line in output.splitlines(): line = line.strip() if not line: continue @@ -326,12 +315,49 @@ def _workflow_tree_at_head( if len(parts) < 3: continue entries.append((parts[2], path)) - if not entries: - return None entries.sort(key=lambda item: item[1]) - files = [path for _sha, path in entries] - content_hash = _fingerprint({"tree": [(sha, path) for sha, path in entries]}) - return files, content_hash + return entries + + +def _workflow_tree_at_ref( + backend: Any, repo_dir: str | None, ref: str +) -> tuple[list[tuple[str, str]], str] | None: + """Return the sorted workflow tree entries and a stable hash for the given ref. + + Returns an empty list when the ref has no workflow files, so callers can detect + additions and deletions against a base ref. + """ + ls_tree = _run_git(backend, repo_dir, f"ls-tree -r {shlex.quote(ref)} -- .github/workflows") + if not ls_tree.ok: + return None + entries = _parse_ls_tree(ls_tree.output) + content_hash = _fingerprint({"tree": entries}) + return entries, content_hash + + +def _fetch_remote_base(backend: Any, repo_dir: str | None, remote: str, branch: str) -> str | None: + """Fetch the base ref from the authenticated remote and return a local alias for it. + + Fetches the pushed branch first; if it does not exist on the remote, fetches the + remote's default branch. Uses `FETCH_HEAD` so the base is bound to the freshly-fetched + remote tip, not a local `refs/remotes/origin/*` ref that the sandbox could rewrite. + """ + fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(remote)} {shlex.quote(branch)}") + if fetch.ok: + return "FETCH_HEAD" + + default = _run_git(backend, repo_dir, "ls-remote --symref origin HEAD") + default_branch = "" + if default.ok: + default_branch = _first_line(default.output).removeprefix("ref: refs/heads/").split("\t")[0] + + for fallback in (default_branch, "dev", "main"): + if not fallback or fallback == branch: + continue + fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(remote)} {shlex.quote(fallback)}") + if fetch.ok: + return "FETCH_HEAD" + return None def _workflow_change_for_push(backend: Any, parsed: ParsedGitPush) -> WorkflowPushChange | None: @@ -354,10 +380,19 @@ def _workflow_change_for_push(backend: Any, parsed: ParsedGitPush) -> WorkflowPu if not head or not _GIT_OBJECT_ID.fullmatch(head): return None - tree = _workflow_tree_at_head(backend, root, head) - if tree is None: + head_tree = _workflow_tree_at_ref(backend, root, head) + if head_tree is None: return None - files, content_hash = tree + + base_ref = _fetch_remote_base(backend, root, parsed.remote, branch_name) + if base_ref is not None: + base_tree = _workflow_tree_at_ref(backend, root, base_ref) + if base_tree is not None and head_tree[0] == base_tree[0]: + # No workflow change against the trusted remote base. + return None + + head_entries, head_content_hash = head_tree + files = sorted({path for _sha, path in head_entries}) remote = _run_git(backend, root, "config --get remote.origin.url") repo = _normalize_remote(_first_line(remote.output)) if remote.ok else "" @@ -371,7 +406,7 @@ def _workflow_change_for_push(backend: Any, parsed: ParsedGitPush) -> WorkflowPu "repo": repo, "branch": branch_name, "files": files, - "content_hash": content_hash, + "content_hash": head_content_hash, } return WorkflowPushChange( fingerprint=_fingerprint(content_payload), diff --git a/tests/test_workflow_push_guard.py b/tests/test_workflow_push_guard.py index 763a25cd..0cbe387c 100644 --- a/tests/test_workflow_push_guard.py +++ b/tests/test_workflow_push_guard.py @@ -25,13 +25,21 @@ class _Backend: *, workflow_files: str = ".github/workflows/ci.yml", tree_entries: dict[str, str] | None = None, + base_tree_entries: dict[str, str] | None = None, + remote_branches: set[str] | None = None, + default_branch: str = "main", ) -> None: self.workflow_files = workflow_files self.tree_entries = tree_entries + # Default base tree is empty so any head workflow file is treated as a change. + self.base_tree_entries = base_tree_entries if base_tree_entries is not None else {} + self.remote_branches = remote_branches or {"feature"} + self.default_branch = default_branch self.commands: list[str] = [] self.head = "a" * 40 + self.base_sha = "b" * 40 - def _ls_tree(self) -> str: + def _head_tree(self) -> str: if self.tree_entries is not None: return "".join( f"100644 blob {sha}\t{path}\n" for path, sha in sorted(self.tree_entries.items()) @@ -39,12 +47,39 @@ class _Backend: files = [path for path in self.workflow_files.split("\n") if path.strip()] return "".join(f"100644 blob {self.head}\t{path}\n" for path in files) + def _base_tree(self) -> str: + if self.base_tree_entries is not None: + return "".join( + f"100644 blob {sha}\t{path}\n" + for path, sha in sorted(self.base_tree_entries.items()) + ) + return self._head_tree() + + def _fetch_branch(self, command: str) -> _Response: + # Format: "git -C /repo fetch origin " or "git fetch origin " + prefix = "fetch origin " + idx = command.find(prefix) + if idx == -1: + return _Response("", 1) + branch = command[idx + len(prefix) :].strip() + if branch in self.remote_branches: + return _Response("") + return _Response("", 1) + def execute(self, command: str, *, timeout: int | None = None) -> _Response: self.commands.append(command) if "rev-parse --show-toplevel" in command: return _Response("/repo\n") if "ls-tree -r" in command: - return _Response(self._ls_tree()) + if "FETCH_HEAD" in command: + return _Response(self._base_tree()) + return _Response(self._head_tree()) + if "fetch origin" in command: + return self._fetch_branch(command) + if "ls-remote --symref origin HEAD" in command: + return _Response( + f"ref: refs/heads/{self.default_branch}\tHEAD\n{self.base_sha}\tHEAD\n" + ) if "config --get remote.origin.url" in command: return _Response("git@github.com:langchain-ai/open-swe.git\n") if "rev-parse --abbrev-ref HEAD" in command: @@ -292,7 +327,13 @@ async def test_workflow_push_restoration_falls_back_when_actions_read_unavailabl async def test_non_workflow_push_runs_without_approval(monkeypatch: pytest.MonkeyPatch) -> None: - guard.SANDBOX_BACKENDS["thread-1"] = _Backend(workflow_files="") + # Head contains workflow files, but they are identical to the trusted remote base, + # so no approval is required for a code-only change. + backend = _Backend( + tree_entries={".github/workflows/ci.yml": "blob-sha-1"}, + base_tree_entries={".github/workflows/ci.yml": "blob-sha-1"}, + ) + guard.SANDBOX_BACKENDS["thread-1"] = backend called = False async def fail_approval(*args: Any, **kwargs: Any) -> bool: @@ -310,6 +351,7 @@ async def test_non_workflow_push_runs_without_approval(monkeypatch: pytest.Monke assert called is True assert isinstance(result, ToolMessage) assert result.content == "pushed" + assert any("fetch origin" in cmd for cmd in backend.commands) async def test_stale_workflow_approval_is_loud_and_blocks(