fix: abort approved workflow push when proxy token elevation fails [closes #97] (#100)
Some checks are pending
CI / Lint (push) Waiting to run
CI / Format check (push) Waiting to run
CI / Unit tests (push) Waiting to run
CI / Playwright E2E (push) Waiting to run
CI / Docker build smoke (push) Waiting to run

* 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:
seahaven-openswe[bot] 2026-07-01 15:05:57 -04:00 • committed by GitHub
parent a52ebed77c
commit 421290d066
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 107 additions and 9 deletions

View file

@ -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
)

View file

@ -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 == []