test(api-server): resolve a streamed chat-completion approval through /v1/runs
This commit is contained in:
@@ -1,23 +1,59 @@
|
||||
"""Legacy /v1/chat/completions streams surface approval requests (#51871)."""
|
||||
"""Legacy /v1/chat/completions streams surface approval requests that the client can resolve
|
||||
through ``POST /v1/runs/{completion_id}/approval`` (#51871)."""
|
||||
import asyncio
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from gateway.platforms.api_server import ThreadSafeAsyncQueue
|
||||
from gateway.platforms.api_server_openai_routes import OpenAICompatRoutesMixin
|
||||
import pytest
|
||||
from aiohttp.test_utils import TestClient, TestServer
|
||||
|
||||
from tests.gateway.test_api_server import _create_app, _make_adapter
|
||||
|
||||
|
||||
class _Host(OpenAICompatRoutesMixin):
|
||||
def __init__(self):
|
||||
self._run_approval_sessions = {}
|
||||
def _approval_agent(decisions):
|
||||
agent = MagicMock()
|
||||
agent.session_prompt_tokens = agent.session_completion_tokens = agent.session_total_tokens = 0
|
||||
|
||||
def run_conversation(**_kw):
|
||||
# The real blocking wait the terminal guard uses, reached through the turn's contextvar.
|
||||
from tools.approval import _gateway_notify_cb
|
||||
from tools.approval_context import get_current_session_key
|
||||
from tools.approval_gateway_wait import _await_gateway_decision
|
||||
key = get_current_session_key(default="")
|
||||
decisions.append(_await_gateway_decision(
|
||||
key, _gateway_notify_cb(key), {"command": "rm -rf /tmp/x", "description": "dangerous",
|
||||
"pattern_key": "rm"}))
|
||||
return {"final_response": "ok", "messages": []}
|
||||
agent.run_conversation.side_effect = run_conversation
|
||||
return agent
|
||||
|
||||
|
||||
def test_stream_approval_keyed_by_completion_id_and_enqueued():
|
||||
async def _go():
|
||||
host, q = _Host(), ThreadSafeAsyncQueue()
|
||||
notify = host._register_stream_approval("chatcmpl-1", q, "sess-1")
|
||||
assert host._run_approval_sessions == {"chatcmpl-1": "chatcmpl-1"}
|
||||
notify({"command": "rm -rf /", "description": "dangerous"})
|
||||
kind, event = await asyncio.wait_for(q.get(), 2)
|
||||
assert kind == "__approval__"
|
||||
assert event["event"] == "approval.request" and event["run_id"] == "chatcmpl-1"
|
||||
assert "once" in event["choices"] and "deny" in event["choices"]
|
||||
asyncio.run(_go())
|
||||
@pytest.mark.asyncio
|
||||
async def test_streamed_completion_approval_resolves_via_runs_endpoint(monkeypatch):
|
||||
# Bound the agent thread's blocking wait so a regression fails fast instead of holding teardown.
|
||||
monkeypatch.setattr("tools.approval_context._get_approval_timeout", lambda: 5)
|
||||
adapter, decisions = _make_adapter(), []
|
||||
app = _create_app(adapter)
|
||||
app.router.add_post("/v1/runs/{run_id}/approval", adapter._handle_run_approval)
|
||||
async with TestClient(TestServer(app)) as cli:
|
||||
with patch.object(adapter, "_create_agent", side_effect=lambda **_k: _approval_agent(decisions)):
|
||||
resp = await cli.post("/v1/chat/completions", json={
|
||||
"model": "hermes-agent", "stream": True, "messages": [{"role": "user", "content": "hi"}]})
|
||||
assert resp.status == 200
|
||||
async def _first_approval():
|
||||
event = None
|
||||
async for raw in resp.content:
|
||||
line = raw.decode().strip()
|
||||
if line == "event: approval.request":
|
||||
event = "next"
|
||||
elif event == "next" and line.startswith("data: "):
|
||||
return json.loads(line[6:])
|
||||
return None
|
||||
event = await asyncio.wait_for(_first_approval(), 10)
|
||||
assert event and event["run_id"].startswith("chatcmpl") and event["command"]
|
||||
approval = await cli.post(f"/v1/runs/{event['run_id']}/approval", json={"choice": "once"})
|
||||
assert approval.status == 200, await approval.text()
|
||||
assert "[DONE]" in (await asyncio.wait_for(resp.text(), 10))
|
||||
assert decisions and decisions[0]["choice"] == "once"
|
||||
assert event["run_id"] not in adapter._run_approval_sessions
|
||||
assert adapter._run_statuses[event["run_id"]]["status"] == "completed"
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Session chat streams surface approval requests keyed by run id (#58856)."""
|
||||
import asyncio
|
||||
from unittest.mock import patch
|
||||
|
||||
from gateway.platforms.api_server import APIServerAdapter, _SessionEventQueue
|
||||
|
||||
@@ -11,7 +12,9 @@ def test_session_stream_approval_keyed_by_run_id_and_enqueued():
|
||||
events = _SessionEventQueue("sess-1", "run_1")
|
||||
notify = host._register_session_stream_approval("run_1", events, "msg_1")
|
||||
assert host._run_approval_sessions == {"run_1": "run_1"}
|
||||
await asyncio.to_thread(notify, {"command": "rm -rf /", "description": "dangerous"})
|
||||
with patch.object(host, "_set_run_status") as set_status:
|
||||
await asyncio.to_thread(notify, {"command": "rm -rf /", "description": "dangerous"})
|
||||
assert set_status.call_args.args == ("run_1", "waiting_for_approval")
|
||||
name, event = await asyncio.wait_for(events.queue.get(), 2)
|
||||
assert name == "approval.request"
|
||||
assert event["run_id"] == "run_1" and event["message_id"] == "msg_1"
|
||||
|
||||
Reference in New Issue
Block a user