diff --git a/agent/middleware/workflow_push_guard.py b/agent/middleware/workflow_push_guard.py index 6d4bd284..d8df8f72 100644 --- a/agent/middleware/workflow_push_guard.py +++ b/agent/middleware/workflow_push_guard.py @@ -44,6 +44,10 @@ _SHELL_OPERATORS = {";", "|", "||", "&"} _REF_NAME = re.compile(r"^[A-Za-z0-9._/@+-]+$") _GIT_OBJECT_ID = re.compile(r"^[0-9a-fA-F]{40,64}$") _UNSAFE_RAW_COMMAND = re.compile(r"[;|`$<>\n\r]") +# Command wrappers that can prefix a `git` invocation (e.g. `command git push`, +# `env FOO=bar git push`, `/usr/bin/git push`). Recognized so a wrapped or path-qualified +# git push cannot slip past the guard as "not a git command". +_GIT_WRAPPERS = {"command", "env", "nice", "sudo", "stdbuf", "nohup", "time", "ionice", "setsid"} class _BlockedGitPush: @@ -177,22 +181,50 @@ def _parse_git_push(command: str) -> ParsedGitPush | _BlockedGitPush | None: return _parse_git_tokens(tokens, repo_dir=None) +def _git_invocation_args(tokens: list[str]) -> list[str] | None: + """Return the args following the `git` executable if these tokens invoke git, else None. + + Recognizes bareword `git`, path invocations like `/usr/bin/git`, and simple command + wrappers (`command`/`env`/`nice`/`sudo`/...), including `env NAME=VALUE ...`. This keeps + a wrapped or path-qualified `git push` from being mistaken for a non-git command and + passed through the guard unchecked. + """ + idx = 0 + while idx < len(tokens): + tok = tokens[idx] + if tok.rsplit("/", 1)[-1] == "git": + return tokens[idx + 1 :] + if tok in _GIT_WRAPPERS: + idx += 1 + # Skip wrapper options and `NAME=VALUE` assignments (e.g. `env FOO=bar git ...`). + while idx < len(tokens) and (tokens[idx].startswith("-") or "=" in tokens[idx]): + idx += 1 + continue + return None + return None + + def _parse_git_tokens( tokens: list[str], *, repo_dir: str | None ) -> ParsedGitPush | _BlockedGitPush | None: - if not tokens or tokens[0] != "git": + git_args = _git_invocation_args(tokens) + if git_args is None: return None - i = 1 - while i < len(tokens) and tokens[i] != "push": - if tokens[i] == "-C" and i + 1 < len(tokens): - repo_dir = tokens[i + 1] + if "push" not in git_args: + # A non-push git command (e.g. `git status`) is not our concern. + return None + i = 0 + while i < len(git_args) and git_args[i] != "push": + if git_args[i] == "-C" and i + 1 < len(git_args): + repo_dir = git_args[i + 1] i += 2 continue - # Any non-push git command (e.g. git status) is not a push, so let it pass. - return None - if i >= len(tokens) or tokens[i] != "push": - return None - return _parse_push_args(tokens[i + 1 :], repo_dir=repo_dir) + # A git push carrying unsupported global options (e.g. `git -c k=v push`, + # `git --git-dir=.git push`). We cannot safely normalize it, so fail closed. + return _BlockedGitPush( + "unsupported git options before `push`; use `git push origin `" + ) + return _parse_push_args(git_args[i + 1 :], repo_dir=repo_dir) def _parse_push_args( @@ -336,17 +368,25 @@ def _workflow_tree_at_ref( 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. + """Fetch the base ref from the exact URL the push will target and return a local alias. - 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. + The base is fetched from the remote's effective *push* URL (`git remote get-url --push`), + not the remote name, because git honors a separate `pushurl`: an untrusted sandbox can + point the fetch URL at attacker/local content while the push still lands on the real + repo. Fetching the base from the push URL keeps the base and the push destination the + same authenticated repo, so a split cannot hide a workflow change. Uses `FETCH_HEAD` + (freshly fetched), never a local `refs/remotes/origin/*` ref the sandbox could rewrite. """ - fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(remote)} {shlex.quote(branch)}") + push_url_res = _run_git(backend, repo_dir, f"remote get-url --push {shlex.quote(remote)}") + push_url = _first_line(push_url_res.output) if push_url_res.ok else "" + if not push_url: + return None + + fetch = _run_git(backend, repo_dir, f"fetch {shlex.quote(push_url)} {shlex.quote(branch)}") if fetch.ok: return "FETCH_HEAD" - default = _run_git(backend, repo_dir, "ls-remote --symref origin HEAD") + default = _run_git(backend, repo_dir, f"ls-remote --symref {shlex.quote(push_url)} HEAD") default_branch = "" if default.ok: default_branch = _first_line(default.output).removeprefix("ref: refs/heads/").split("\t")[0] @@ -354,7 +394,9 @@ def _fetch_remote_base(backend: Any, repo_dir: str | None, remote: str, branch: 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)}") + fetch = _run_git( + backend, repo_dir, f"fetch {shlex.quote(push_url)} {shlex.quote(fallback)}" + ) if fetch.ok: return "FETCH_HEAD" return None @@ -382,7 +424,22 @@ def _workflow_change_for_push(backend: Any, parsed: ParsedGitPush) -> WorkflowPu head_tree = _workflow_tree_at_ref(backend, root, head) if head_tree is None: - return None + # The workflow tree at head could not be read (an empty tree returns ([], hash), so + # None means the `ls-tree` read genuinely failed). Fail closed rather than skipping + # the guard, mirroring the base-read path. + return WorkflowPushChange( + fingerprint="", + repo="", + branch=branch_name, + files=[], + head_sha=head, + remote=parsed.remote, + local_ref=parsed.local_ref, + remote_ref=parsed.remote_ref, + fixed_command="", + blocked=True, + blocked_reason="could not read the workflow tree at the pushed head; blocking to be safe", + ) base_ref = _fetch_remote_base(backend, root, parsed.remote, branch_name) if base_ref is not None: diff --git a/tests/test_workflow_push_guard.py b/tests/test_workflow_push_guard.py index 0cbe387c..6edc93f8 100644 --- a/tests/test_workflow_push_guard.py +++ b/tests/test_workflow_push_guard.py @@ -28,6 +28,8 @@ class _Backend: base_tree_entries: dict[str, str] | None = None, remote_branches: set[str] | None = None, default_branch: str = "main", + push_url: str = "https://github.com/langchain-ai/open-swe.git", + ls_tree_head_ok: bool = True, ) -> None: self.workflow_files = workflow_files self.tree_entries = tree_entries @@ -35,6 +37,8 @@ class _Backend: 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.push_url = push_url + self.ls_tree_head_ok = ls_tree_head_ok self.commands: list[str] = [] self.head = "a" * 40 self.base_sha = "b" * 40 @@ -56,12 +60,14 @@ class _Backend: 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) + # Format: "git [-C /repo] fetch " + idx = command.find(" fetch ") if idx == -1: return _Response("", 1) - branch = command[idx + len(prefix) :].strip() + parts = command[idx + len(" fetch ") :].split() + if len(parts) < 2: + return _Response("", 1) + branch = parts[-1] if branch in self.remote_branches: return _Response("") return _Response("", 1) @@ -70,13 +76,15 @@ class _Backend: self.commands.append(command) if "rev-parse --show-toplevel" in command: return _Response("/repo\n") + if "remote get-url --push" in command: + return _Response("", 1) if not self.push_url else _Response(f"{self.push_url}\n") if "ls-tree -r" in command: if "FETCH_HEAD" in command: return _Response(self._base_tree()) - return _Response(self._head_tree()) - if "fetch origin" in command: + return _Response(self._head_tree()) if self.ls_tree_head_ok else _Response("", 1) + if " fetch " in command: return self._fetch_branch(command) - if "ls-remote --symref origin HEAD" in command: + if "ls-remote --symref" in command: return _Response( f"ref: refs/heads/{self.default_branch}\tHEAD\n{self.base_sha}\tHEAD\n" ) @@ -351,7 +359,9 @@ 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) + # The base is fetched from the effective push URL, not the `origin` remote name. + assert any("remote get-url --push" in cmd for cmd in backend.commands) + assert any(f" fetch {backend.push_url}" in cmd for cmd in backend.commands) async def test_stale_workflow_approval_is_loud_and_blocks( @@ -687,3 +697,106 @@ async def test_approved_workflow_push_runs_on_non_langsmith_providers( assert isinstance(result, ToolMessage) assert result.content == "pushed" assert refresh_calls == [] + + +def test_parse_git_push_guards_wrapped_and_optioned_forms() -> None: + # Path-qualified and wrapped git invocations must still be recognized as pushes + # (so the guard engages), not passed through as "not a git command". + for command in ( + "/usr/bin/git push origin feature", + "command git push origin feature", + "env FOO=bar git push origin feature", + ): + parsed = guard._parse_git_push(command) + assert isinstance(parsed, guard.ParsedGitPush), command + assert parsed.remote_ref == "feature" + + # A push carrying unsupported global options cannot be normalized, so it is blocked + # (fail closed) rather than run unguarded. + for command in ( + "git -c protocol.version=2 push origin feature", + "git --git-dir=.git push origin feature", + "git --no-pager push origin feature", + ): + assert isinstance(guard._parse_git_push(command), guard._BlockedGitPush), command + + # A non-push git command is still ignored. + assert guard._parse_git_push("/usr/bin/git status") is None + + +async def test_base_fetched_from_push_url_not_origin_remote( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # A changed workflow file at head (vs the trusted base fetched from the push URL) must + # require approval, and the base must be fetched from the push URL, never the bare + # `origin` remote name (which the sandbox can repoint via `remote set-url --push`). + backend = _Backend( + tree_entries={".github/workflows/ci.yml": "head-sha"}, + base_tree_entries={".github/workflows/ci.yml": "base-sha"}, + ) + guard.SANDBOX_BACKENDS["thread-1"] = backend + + async def fake_approved(thread_id: str, fingerprint: str) -> bool: + return False + + async def fake_rejected(thread_id: str, fingerprint: str) -> bool: + return False + + async def fake_find_approval(*args: Any, **kwargs: Any) -> dict[str, Any] | None: + return None + + async def fake_pending(thread_id: str, **kwargs: Any) -> tuple[dict[str, Any], bool]: + return {"fingerprint": kwargs["fingerprint"], "status": "pending", "notified": False}, True + + async def fake_post(*args: Any, **kwargs: Any) -> tuple[str, None]: + return "1700000000.000300", None + + async def fake_notified(*args: Any, **kwargs: Any) -> None: + return None + + monkeypatch.setattr(guard, "workflow_push_approved", fake_approved) + monkeypatch.setattr(guard, "workflow_push_rejected", fake_rejected) + monkeypatch.setattr(guard, "find_workflow_push_approval", fake_find_approval) + monkeypatch.setattr(guard, "ensure_workflow_push_pending", fake_pending) + monkeypatch.setattr(guard, "post_slack_thread_reply_with_ts", fake_post) + monkeypatch.setattr(guard, "mark_workflow_push_notified", fake_notified) + + async def handler(_request: Any) -> ToolMessage: + return ToolMessage(content="pushed", tool_call_id="call-1") + + result = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) + + assert isinstance(result, ToolMessage) + assert result.status == "error" + payload = json.loads(str(result.content)) + assert payload["workflow_approval_status"] == "approval_required" + assert any("remote get-url --push" in cmd for cmd in backend.commands) + assert any(f" fetch {backend.push_url}" in cmd for cmd in backend.commands) + assert not any(" fetch origin " in cmd for cmd in backend.commands) + + +async def test_unreadable_head_tree_blocks(monkeypatch: pytest.MonkeyPatch) -> None: + # If `ls-tree` on the head cannot be read, the guard must fail closed (block), not + # skip the guard and run the original push. + backend = _Backend(ls_tree_head_ok=False) + guard.SANDBOX_BACKENDS["thread-1"] = backend + + async def fail_approval(*args: Any, **kwargs: Any) -> bool: + raise AssertionError("approval should not be checked for a blocked push") + + monkeypatch.setattr(guard, "workflow_push_approved", fail_approval) + + called = False + + async def handler(_request: Any) -> ToolMessage: + nonlocal called + called = True + return ToolMessage(content="pushed", tool_call_id="call-1") + + result = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) + + assert called is False + assert isinstance(result, ToolMessage) + assert result.status == "error" + payload = json.loads(str(result.content)) + assert payload["error_type"] == "WorkflowPushBlocked"