mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 16:13:15 +00:00
* Abort approved workflow pushes when proxy elevation fails When _run_with_workflow_token cannot elevate the sandbox proxy token to workflows:write, it now returns a clear ToolMessage error instead of running the push over the base token and getting a raw GitHub remote rejection. Refs: #97 * Fix workflow push guard crash and non-langsmith regression - Thread the ToolCallRequest into _run_with_workflow_token so the WorkflowPushElevationFailed ToolMessage is stamped with the real tool_call_id instead of an empty id that crashes the Anthropic API. - Only perform the elevation/abort path on SANDBOX_TYPE=langsmith; other providers run the approved push directly. - Add tests covering the real tool_call_id, single refresh call, and non-langsmith approved pushes. Refs: #97 --------- Co-authored-by: amoussa1229 <166072409+amoussa1229@users.noreply.github.com>
This commit is contained in:
parent
a52ebed77c
commit
421290d066
2 changed files with 107 additions and 9 deletions
|
|
@ -6,6 +6,7 @@ import asyncio
|
|||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import shlex
|
||||
import threading
|
||||
|
|
@ -472,22 +473,44 @@ async def _approval_state(request: ToolCallRequest, change: WorkflowPushChange)
|
|||
|
||||
async def _run_with_workflow_token(
|
||||
thread_id: str,
|
||||
request: ToolCallRequest,
|
||||
run: Callable[[], Awaitable[ToolMessage | Command]],
|
||||
) -> ToolMessage | Command:
|
||||
sandbox_type = os.getenv("SANDBOX_TYPE", "langsmith")
|
||||
if sandbox_type != "langsmith":
|
||||
return await run()
|
||||
|
||||
elevated = await refresh_proxy_token(
|
||||
thread_id, permissions=WORKFLOW_RUNTIME_PROXY_TOKEN_PERMISSIONS
|
||||
)
|
||||
if not elevated:
|
||||
logger.error(
|
||||
"Workflow push approved for thread %s, but proxy token elevation to workflows:write "
|
||||
"failed; the sandbox cannot push workflow files without an elevated token.",
|
||||
thread_id,
|
||||
)
|
||||
error_message = ToolMessage(
|
||||
content=json.dumps(
|
||||
{
|
||||
"status": "error",
|
||||
"error_type": "WorkflowPushElevationFailed",
|
||||
"error": (
|
||||
"Workflow push approved, but the sandbox could not obtain a "
|
||||
"workflows-scoped token. Please retry the push or check the "
|
||||
"GitHub proxy / token minting configuration."
|
||||
),
|
||||
}
|
||||
),
|
||||
tool_call_id="",
|
||||
status="error",
|
||||
)
|
||||
return _tool_message_for_request(error_message, request)
|
||||
try:
|
||||
return await run()
|
||||
finally:
|
||||
if elevated:
|
||||
restored = await refresh_proxy_token(
|
||||
thread_id, permissions=RUNTIME_PROXY_TOKEN_PERMISSIONS
|
||||
)
|
||||
if not restored:
|
||||
await refresh_proxy_token(
|
||||
thread_id, permissions=BASE_RUNTIME_PROXY_TOKEN_PERMISSIONS
|
||||
)
|
||||
restored = await refresh_proxy_token(thread_id, permissions=RUNTIME_PROXY_TOKEN_PERMISSIONS)
|
||||
if not restored:
|
||||
await refresh_proxy_token(thread_id, permissions=BASE_RUNTIME_PROXY_TOKEN_PERMISSIONS)
|
||||
|
||||
|
||||
class WorkflowPushGuardMiddleware(AgentMiddleware):
|
||||
|
|
@ -519,7 +542,7 @@ class WorkflowPushGuardMiddleware(AgentMiddleware):
|
|||
state = await _approval_state(request, change)
|
||||
if state == "approved" and thread_id:
|
||||
safe_request = _override_execute_command(request, change.fixed_command)
|
||||
return await _run_with_workflow_token(thread_id, lambda: handler(safe_request))
|
||||
return await _run_with_workflow_token(thread_id, request, lambda: handler(safe_request))
|
||||
return _tool_message_for_request(
|
||||
_blocked_message(change, already_rejected=state == "rejected"), request
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
|
@ -273,3 +274,77 @@ async def test_non_workflow_push_runs_without_approval(monkeypatch: pytest.Monke
|
|||
assert called is True
|
||||
assert isinstance(result, ToolMessage)
|
||||
assert result.content == "pushed"
|
||||
|
||||
|
||||
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 == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue