from __future__ import annotations import json import os from typing import Any import pytest from langchain_core.messages import ToolMessage from agent.middleware import workflow_push_guard as guard class _Response: def __init__(self, output: str, exit_code: int = 0) -> None: self.output = output self.exit_code = exit_code self.truncated = False class _Backend: id = "sandbox-id" def __init__( self, *, 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 _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()) ) 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: 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: return _Response("feature\n") if "rev-parse HEAD" in command or "rev-parse feature" in command: return _Response(f"{self.head}\n") return _Response("") class _Runtime: config = { "configurable": { "thread_id": "thread-1", "slack_thread": {"channel_id": "C123", "thread_ts": "1700000000.000100"}, } } class _Request: runtime = _Runtime() def __init__(self, command: str = "git -C /repo push origin feature") -> None: self.tool_call = { "name": "execute", "args": {"command": command}, "id": "call-1", } def override(self, **kwargs: Any) -> _Request: next_request = _Request() next_request.tool_call = kwargs.get("tool_call", self.tool_call) return next_request @pytest.fixture(autouse=True) def _clear_backend_cache() -> Any: guard.SANDBOX_BACKENDS.clear() yield guard.SANDBOX_BACKENDS.clear() def test_parse_git_push_supports_git_c_and_cd() -> None: assert guard._parse_git_push("git -C /repo push origin feature") == guard.ParsedGitPush( repo_dir="/repo", remote="origin", local_ref="feature", remote_ref="feature" ) assert guard._parse_git_push( "cd /repo && git push -u origin HEAD:feature" ) == guard.ParsedGitPush( repo_dir="/repo", remote="origin", local_ref="HEAD", remote_ref="feature", set_upstream=True, ) def test_parse_git_push_blocks_unsafe_and_unrecognized_forms() -> None: assert isinstance(guard._parse_git_push("git status && git push"), guard._BlockedGitPush) assert isinstance( guard._parse_git_push("git push origin feature; git push origin evil:feature"), guard._BlockedGitPush, ) assert isinstance( guard._parse_git_push("git push --force origin feature"), guard._BlockedGitPush ) assert isinstance(guard._parse_git_push("git push origin"), guard._BlockedGitPush) assert isinstance( guard._parse_git_push("git push origin HEAD~1:feature"), guard._BlockedGitPush ) assert guard._parse_git_push("git status") is None def test_workflow_change_for_push_fingerprints_head_workflow_tree() -> None: backend = _Backend() change = guard._workflow_change_for_push( backend, guard.ParsedGitPush( repo_dir="/repo", remote="origin", local_ref="feature", remote_ref="feature" ), ) assert change is not None assert change.repo == "https://github.com/langchain-ai/open-swe" assert change.branch == "feature" assert change.files == [".github/workflows/ci.yml"] assert ( change.fixed_command == "git -C /repo push origin aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa:refs/heads/feature" ) assert len(change.fingerprint) == 64 def test_workflow_change_for_push_ignores_non_workflow_push() -> None: backend = _Backend(workflow_files="") assert ( guard._workflow_change_for_push( backend, guard.ParsedGitPush( repo_dir="/repo", remote="origin", local_ref="feature", remote_ref="feature" ), ) is None ) def test_workflow_change_for_push_rejects_non_current_refspec() -> None: backend = _Backend() assert ( guard._workflow_change_for_push( backend, guard.ParsedGitPush( repo_dir="/repo", remote="origin", local_ref="evil", remote_ref="feature" ), ) is None ) async def test_unapproved_workflow_push_blocks_and_posts_slack( monkeypatch: pytest.MonkeyPatch, ) -> None: guard.SANDBOX_BACKENDS["thread-1"] = _Backend() posted: dict[str, Any] = {} 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( channel_id: str, thread_ts: str, message: str, **kwargs: Any ) -> tuple[str, None]: posted.update( channel_id=channel_id, thread_ts=thread_ts, message=message, blocks=kwargs["blocks"] ) return "1700000000.000200", None async def fake_notified(thread_id: str, fingerprint: str) -> None: posted["notified"] = fingerprint 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) 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["workflow_approval_status"] == "approval_required" assert payload["files"] == [".github/workflows/ci.yml"] assert posted["channel_id"] == "C123" assert posted["blocks"][1]["elements"][0]["value"] async def test_approved_workflow_push_elevates_and_restores( monkeypatch: pytest.MonkeyPatch, ) -> None: guard.SANDBOX_BACKENDS["thread-1"] = _Backend() refreshed: list[dict[str, str]] = [] async def fake_approved(thread_id: str, fingerprint: str) -> bool: return True async def fake_refresh(thread_id: str | None, *, permissions: dict[str, str]) -> bool: refreshed.append(dict(permissions)) return True async def fake_find_approval(*args: Any, **kwargs: Any) -> dict[str, Any] | None: return None monkeypatch.setattr(guard, "workflow_push_approved", fake_approved) monkeypatch.setattr(guard, "find_workflow_push_approval", fake_find_approval) monkeypatch.setattr(guard, "refresh_proxy_token", fake_refresh) pushed_command = "" async def handler(request: Any) -> ToolMessage: nonlocal pushed_command pushed_command = request.tool_call["args"]["command"] return ToolMessage(content="pushed", tool_call_id="call-1") result = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) assert isinstance(result, ToolMessage) assert result.content == "pushed" assert ( pushed_command == "git -C /repo push origin aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa:refs/heads/feature" ) assert refreshed[0]["workflows"] == "write" assert "workflows" not in refreshed[1] assert refreshed[1]["actions"] == "read" async def test_workflow_push_restoration_falls_back_when_actions_read_unavailable( monkeypatch: pytest.MonkeyPatch, ) -> None: guard.SANDBOX_BACKENDS["thread-1"] = _Backend() refreshed: list[dict[str, str]] = [] async def fake_approved(thread_id: str, fingerprint: str) -> bool: return True async def fake_refresh(thread_id: str | None, *, permissions: dict[str, str]) -> bool: refreshed.append(dict(permissions)) return "actions" not in permissions async def fake_find_approval(*args: Any, **kwargs: Any) -> dict[str, Any] | None: return None monkeypatch.setattr(guard, "workflow_push_approved", fake_approved) monkeypatch.setattr(guard, "find_workflow_push_approval", fake_find_approval) monkeypatch.setattr(guard, "refresh_proxy_token", fake_refresh) async def handler(_request: Any) -> ToolMessage: return ToolMessage(content="pushed", tool_call_id="call-1") await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) assert refreshed[0]["workflows"] == "write" assert refreshed[1]["actions"] == "read" assert refreshed[2] == guard.BASE_RUNTIME_PROXY_TOKEN_PERMISSIONS assert "actions" not in refreshed[2] async def test_non_workflow_push_runs_without_approval(monkeypatch: pytest.MonkeyPatch) -> None: # 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: raise AssertionError("approval should not be checked") monkeypatch.setattr(guard, "workflow_push_approved", fail_approval) 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 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( monkeypatch: pytest.MonkeyPatch, ) -> None: guard.SANDBOX_BACKENDS["thread-1"] = _Backend() posted: dict[str, Any] = {} 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 {"fingerprint": "old-fp", "status": "approved", "decided_at": "2024-01-01T00:00:00"} 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( channel_id: str, thread_ts: str, message: str, **kwargs: Any ) -> tuple[str, None]: posted.update( channel_id=channel_id, thread_ts=thread_ts, message=message, blocks=kwargs["blocks"] ) return "1700000000.000300", None async def fake_notified(thread_id: str, fingerprint: str) -> None: posted["notified"] = fingerprint 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 "changed since that approval" in payload["error"] assert posted["channel_id"] == "C123" async def test_rebased_workflow_push_keeps_fingerprint_for_same_tree( monkeypatch: pytest.MonkeyPatch, ) -> None: # Fingerprint is based on the workflow tree at head, not on a diff against origin. backend = _Backend( tree_entries={".github/workflows/ci.yml": "blob-sha-1"}, ) backend.head = "b" * 40 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["fingerprint"] assert payload["files"] == [".github/workflows/ci.yml"] # Rebase with a new head but the same workflow tree -> fingerprint stays stable. backend.head = "c" * 40 result2 = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) assert isinstance(result2, ToolMessage) payload2 = json.loads(str(result2.content)) assert payload2["fingerprint"] == payload["fingerprint"] # A content change to the same workflow file produces a different fingerprint. backend.tree_entries = {".github/workflows/ci.yml": "blob-sha-2"} result3 = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) assert isinstance(result3, ToolMessage) payload3 = json.loads(str(result3.content)) assert payload3["fingerprint"] != payload["fingerprint"] async def test_deleted_workflow_file_requires_approval( monkeypatch: pytest.MonkeyPatch, ) -> None: # Deleting a workflow file means it is no longer in the head workflow tree. The # guard should still trigger if there are other workflow files at head; if the last # workflow file is deleted, the push is no longer workflow-guarded. backend = _Backend( tree_entries={ ".github/workflows/ci.yml": "blob-sha-1", ".github/workflows/other.yml": "blob-sha-2", }, ) 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") # Push with the full workflow tree present requires approval. 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 ".github/workflows/ci.yml" in payload["files"] # A push whose head tree has deleted ci.yml but still contains other.yml still # requires approval, with a different fingerprint. backend.tree_entries = {".github/workflows/other.yml": "blob-sha-2"} result2 = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) assert isinstance(result2, ToolMessage) payload2 = json.loads(str(result2.content)) assert payload2["workflow_approval_status"] == "approval_required" assert payload2["files"] == [".github/workflows/other.yml"] assert payload2["fingerprint"] != payload["fingerprint"] async def test_base_poisoning_does_not_bypass_guard( monkeypatch: pytest.MonkeyPatch, ) -> None: # The sandbox can rewrite refs/remotes/origin/*; the guard must use the workflow # tree at the pushed head, not a diff against a remote ref. This backend records # whether any command touches the remote-tracking ref. backend = _Backend( tree_entries={ ".github/workflows/ci.yml": "blob-sha-1", ".github/workflows/evil.yml": "blob-sha-evil", }, ) 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" # The evil workflow file present at head must appear in the approval list. assert ".github/workflows/evil.yml" in payload["files"] # No remote-tracking ref should have been consulted. assert not any("refs/remotes/origin" in cmd for cmd in backend.commands) async def test_exact_rejection_checked_before_stale_approval( monkeypatch: pytest.MonkeyPatch, ) -> None: backend = _Backend() 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 True # A prior approval for the same repo/branch/files exists; the stale path would # normally match. But because the exact fingerprint is rejected, we must return # "rejected", not "stale_approval". async def fake_find_approval(*args: Any, **kwargs: Any) -> dict[str, Any] | None: return {"fingerprint": "old-fp", "status": "approved", "decided_at": "2024-01-01T00:00:00"} async def fake_pending(*args: Any, **kwargs: Any) -> tuple[dict[str, Any], bool]: raise AssertionError("should not create a new pending record") 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) 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"] == "rejected" async def test_approved_workflow_push_aborts_when_elevation_fails( monkeypatch: pytest.MonkeyPatch, ) -> None: guard.SANDBOX_BACKENDS["thread-1"] = _Backend() async def fake_approved(thread_id: str, fingerprint: str) -> bool: return True refresh_calls: list[dict[str, str]] = [] async def fake_refresh(*args: Any, **kwargs: Any) -> bool: refresh_calls.append(dict(kwargs.get("permissions", {}))) return False monkeypatch.setattr(guard, "workflow_push_approved", fake_approved) monkeypatch.setattr(guard, "refresh_proxy_token", fake_refresh) called = False request = _Request() 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.tool_call_id == "call-1" assert result.status == "error" payload = json.loads(str(result.content)) assert payload["status"] == "error" assert payload["error_type"] == "WorkflowPushElevationFailed" assert "workflows-scoped token" in payload["error"] assert len(refresh_calls) == 1 assert refresh_calls[0].get("workflows") == "write" async def test_approved_workflow_push_runs_on_non_langsmith_providers( monkeypatch: pytest.MonkeyPatch, ) -> None: guard.SANDBOX_BACKENDS["thread-1"] = _Backend() async def fake_approved(thread_id: str, fingerprint: str) -> bool: return True refresh_calls: list[dict[str, str]] = [] async def fake_refresh(*args: Any, **kwargs: Any) -> bool: refresh_calls.append(dict(kwargs.get("permissions", {}))) return False monkeypatch.setattr(guard, "workflow_push_approved", fake_approved) monkeypatch.setattr(guard, "refresh_proxy_token", fake_refresh) monkeypatch.setattr(guard, "os", os) called = False async def handler(request: Any) -> ToolMessage: nonlocal called called = True return ToolMessage(content="pushed", tool_call_id=request.tool_call["id"]) with monkeypatch.context() as mp: mp.setenv("SANDBOX_TYPE", "local") result = await guard.WorkflowPushGuardMiddleware().awrap_tool_call(_Request(), handler) assert called is True assert isinstance(result, ToolMessage) assert result.content == "pushed" assert refresh_calls == []