From ee8eeb31720c469e9d59609ed6574ffc6d15932e Mon Sep 17 00:00:00 2001 From: liuhao1024 Date: Thu, 24 Sep 2026 16:46:45 +0530 Subject: [PATCH] fix(api-server): emit approval events on legacy chat-completions SSE stream Hand-grafted from #51878 onto the split api_server_openai_routes.py; approvals keyed by completion id and resolved via POST /v1/runs/{id}/approval (#51871). --- gateway/platforms/api_server.py | 26 +++++++++++++-- gateway/platforms/api_server_openai_routes.py | 32 ++++++++++++++++++- .../gateway/test_chat_completions_approval.py | 23 +++++++++++++ 3 files changed, 77 insertions(+), 4 deletions(-) create mode 100644 tests/gateway/test_chat_completions_approval.py diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index a418ee5625..9bc91cf94c 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -3982,8 +3982,11 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter): confirmed_runtime_lock: bool = False, bind_declared_conversation: bool = False, session_history_delivery: str = "", turn_author: Optional[Dict[str, Any]] = None, relay_metadata: Optional[Dict[str, Any]] = None, notification_category: str = "result", - resume_unanswered_turn: bool = False) -> tuple: + resume_unanswered_turn: bool = False, approval_notify_callback=None, + approval_session_key: Optional[str] = None) -> tuple: """Create an agent and run one turn in a thread executor -> ``(result, usage)``. + ``approval_notify_callback`` (with ``approval_session_key``) routes dangerous-command + approval requests to the caller's stream, keyed like ``/v1/runs`` approvals (#51871). ``agent_ref[0]`` receives the agent so SSE writers can interrupt it; ``active_run_id`` registers it in ``_active_run_agents``. Under a confirmed model lock the actual provider/model must match or the turn fails; ``runtime`` metadata is attached. @@ -4060,8 +4063,25 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter): ) if relay_metadata: conversation_kwargs["relay_metadata"] = relay_metadata - with notification_turn(agent, muted=muted, session_id=session_id or ""): - result = agent.run_conversation(**conversation_kwargs) + approval_token = None + if approval_notify_callback is not None and approval_session_key: + # Same machinery as /v1/runs (_run_agent_sync): the contextvar scopes + # this turn's approvals to the key the resolve endpoint looks up. + from tools.approval import register_gateway_notify + from tools.approval_context import set_current_session_key + approval_token = set_current_session_key(approval_session_key) + register_gateway_notify(approval_session_key, approval_notify_callback) + try: + with notification_turn(agent, muted=muted, session_id=session_id or ""): + result = agent.run_conversation(**conversation_kwargs) + finally: + if approval_token is not None: + from tools.approval import unregister_gateway_notify + from tools.approval_context import reset_current_session_key + with suppress(Exception): + unregister_gateway_notify(approval_session_key) + with suppress(Exception): + reset_current_session_key(approval_token) result, usage = self._finish_turn_result( agent, result, session_id, route=route, requested_runtime=requested_runtime, route_source=route_source, confirmed_runtime_lock=confirmed_runtime_lock) diff --git a/gateway/platforms/api_server_openai_routes.py b/gateway/platforms/api_server_openai_routes.py index c64d1b7134..4ed8177b06 100644 --- a/gateway/platforms/api_server_openai_routes.py +++ b/gateway/platforms/api_server_openai_routes.py @@ -538,6 +538,28 @@ class OpenAICompatRoutesMixin: requested_provider=overrides.get("requested_provider"), route=route) return route, overrides, (_error_response(err, 400) if err else None) + def _register_stream_approval(self, completion_id, stream_q, session_id): + """Expose a streaming completion to ``POST /v1/runs/{completion_id}/approval`` and return + the notify callback that pushes ``approval.request`` onto ``stream_q`` (#51871). Keyed by + the completion id (never the shared session key) so concurrent turns can't cross-resolve.""" + from gateway.platforms.api_server import _approval_event_choices + self._run_approval_sessions[completion_id] = completion_id + + def _approval_notify(approval_data): + event = dict(approval_data or {}) + if "command" in event: + from gateway.run import _redact_approval_command + event["command"] = _redact_approval_command(event.get("command")) + event.update({ + "event": "approval.request", "run_id": completion_id, "session_id": session_id or "", + "timestamp": time.time(), + "choices": _approval_event_choices( + smart_denied=bool(event.get("smart_denied")), + allow_session=event.get("allow_session") is not False, + allow_permanent=event.get("allow_permanent") is not False)}) + stream_q.put_threadsafe(("__approval__", event)) + return _approval_notify + def _spawn_stream_agent(self, stream_q, **run_kwargs) -> tuple: """Start ``_run_agent`` for an SSE writer -> ``(agent_task, agent_ref)``. ``agent_ref[0]`` lets the writer interrupt on disconnect; the EOS sentinel is enqueued from the task's done @@ -702,9 +724,15 @@ class OpenAICompatRoutesMixin: # tool_progress_callback deliberately NOT wired: it would duplicate the structured # start/complete callbacks (which carry the tool_call id). + approval_notify = self._register_stream_approval(completion_id, _stream_q, session_id) agent_task, agent_ref = self._spawn_stream_agent( _stream_q, tool_start_callback=_on_tool_start, - tool_complete_callback=_on_tool_complete, **run_kwargs) + tool_complete_callback=_on_tool_complete, approval_notify_callback=approval_notify, + approval_session_key=completion_id, **run_kwargs) + # The completion id doubles as the run id: drop the approval mapping once the turn + # ends so POST /v1/runs/{id}/approval answers 409 rather than resolving stale keys. + agent_task.add_done_callback( + lambda _fut: self._run_approval_sessions.pop(completion_id, None)) # #13437 identity contract: an explicit-header client keeps addressing the id it # sent; the response echoes that stable id while reads/writes adopt the live tip, # so a rotation mid-turn (after these headers are prepared) never changes what the @@ -844,6 +872,8 @@ class OpenAICompatRoutesMixin: await response.write(_sse_frame(_chunk({"reasoning_content": delta[1]}))) elif isinstance(delta, tuple) and len(delta) == 2 and delta[0] == "__status__": await response.write(_sse_frame(delta[1], event="hermes.status")) + elif isinstance(delta, tuple) and len(delta) == 2 and delta[0] == "__approval__": + await response.write(_sse_frame(delta[1], event="approval.request")) else: await response.write(_sse_frame(_chunk({"content": delta}))) # The agent can fail after the queue drains (task raises / result flagged failed or diff --git a/tests/gateway/test_chat_completions_approval.py b/tests/gateway/test_chat_completions_approval.py new file mode 100644 index 0000000000..44d0eaf5f7 --- /dev/null +++ b/tests/gateway/test_chat_completions_approval.py @@ -0,0 +1,23 @@ +"""Legacy /v1/chat/completions streams surface approval requests (#51871).""" +import asyncio + +from gateway.platforms.api_server import ThreadSafeAsyncQueue +from gateway.platforms.api_server_openai_routes import OpenAICompatRoutesMixin + + +class _Host(OpenAICompatRoutesMixin): + def __init__(self): + self._run_approval_sessions = {} + + +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())