fix(api-server): deliver transform_llm_output rewrites on streaming routes

This commit is contained in:
Brooklyn Nicholson
2026-09-24 23:16:35 -05:00
committed by brooklyn!
parent 7fa45eb349
commit e7b2db974c
2 changed files with 116 additions and 1 deletions

View File

@@ -82,6 +82,22 @@ def _hermes_extras(completed, is_partial, is_failed, err_msg, finish_reason: str
"error_code": "output_truncated" if finish_reason == "length" else "agent_error"}
_TRANSFORMED_NOTICE = "\n\n[Response transformed after streaming]\n"
def _post_stream_transform(result: Any) -> tuple:
"""``(text, appended)`` a stream still owes after ``transform_llm_output`` rewrote the final
(the deltas already carried the raw reply): the appended suffix, or the whole rewrite with
``appended=False`` when it is not a pure append. ``("", False)`` when nothing was transformed."""
if not isinstance(result, dict) or not result.get("response_transformed"):
return "", False
final = result.get("final_response") or ""
original = result.get("pre_transform_response") or ""
if original and final.startswith(original):
return final[len(original):], True
return final, False
def _message_item(text: Any) -> Dict[str, Any]:
"""Responses ``message`` output item carrying one ``output_text`` part."""
return {"type": "message", "role": "assistant",
@@ -219,6 +235,7 @@ class _ResponsesStream:
self.message_opened = False
self.reasoning_item: Optional[Dict[str, Any]] = None # open ``reasoning`` output item
self.final_response_text = ""
self.transformed_final = "" # non-append transform_llm_output rewrite; replaces the deltas
self.agent_error: Optional[str] = None
self.usage: Dict[str, int] = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
self.terminal_snapshot_persisted = False
@@ -456,6 +473,12 @@ class _ResponsesStream:
self.result = result
self.usage = agent_usage or self.usage
agent_final = result.get("final_response", "") if isinstance(result, dict) else ""
tail, appended = _post_stream_transform(result)
if tail and self.final_text_parts:
if appended:
await self.emit_text_delta(tail)
else:
self.transformed_final = agent_final
if agent_final and not self.final_text_parts:
await self.emit_text_delta(agent_final)
if agent_final and not self.final_response_text:
@@ -468,7 +491,8 @@ class _ResponsesStream:
async def close_message_item(self) -> None:
await self.close_reasoning_item()
self.final_response_text = "".join(self.final_text_parts) or self.final_response_text
self.final_response_text = (
self.transformed_final or "".join(self.final_text_parts) or self.final_response_text)
if not self.message_opened:
return
await self.write_event("response.output_text.done", {
@@ -921,6 +945,11 @@ class OpenAICompatRoutesMixin:
fallback_text = _resolve_media_to_data_urls(result.get("final_response") or "")
if fallback_text:
await response.write(_sse_frame(_chunk({"content": fallback_text})))
elif not presentation_muted:
# Chat chunks can only append: a non-append rewrite follows the streamed text (as in the CLI).
tail, appended = _post_stream_transform(result)
if tail:
await response.write(_sse_frame(_chunk({"content": tail if appended else _TRANSFORMED_NOTICE + tail})))
if finish_reason != "stop":
if err_msg and not presentation_muted:
finish_chunk["error"] = {

View File

@@ -0,0 +1,86 @@
"""``transform_llm_output`` on the api_server streaming routes (#119323).
The deltas carry the raw model reply; the hook rewrites the final afterwards. Streaming
``/v1/chat/completions`` and ``/v1/responses`` must still deliver the rewrite, like their
non-streaming twins do.
"""
import json
from contextlib import ExitStack
from unittest.mock import patch
import pytest
from aiohttp import web
from aiohttp.test_utils import TestClient, TestServer
from agent.turn_finalizer import apply_llm_output_transform
from gateway.config import PlatformConfig
from gateway.platforms.api_server import APIServerAdapter
RAW = "Original model reply."
TRANSFORMS = {
"append": lambda text: f"{text}\n\n[PLUGIN footer]",
"replace": lambda text: f"[PLUGIN warning] {text}",
}
def _app_with_fake_turn(adapter, transform, stack):
class FakeAgent:
def __init__(self, cb):
self.cb, self.session_id, self.model, self.platform = cb, "s", "fake", "api_server"
def run_conversation(self, user_message, conversation_history, task_id=None, **kw):
for part in ("Original ", "model ", "reply."):
self.cb(part)
final, transformed, pre = apply_llm_output_transform(self, RAW, turn_id="t1")
return {"final_response": final, "response_transformed": transformed,
"pre_transform_response": pre, "messages": [], "api_calls": 1, "completed": True}
def __getattr__(self, name):
return None
def invoke_hook(name, **kw):
return [transform(kw["response_text"])] if name == "transform_llm_output" else []
app = web.Application()
app.router.add_post("/v1/chat/completions", adapter._handle_chat_completions)
app.router.add_post("/v1/responses", adapter._handle_responses)
finish = lambda agent, result, session_id, **kw: (result, {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2})
stack.enter_context(patch("hermes_cli.lifecycle.invoke_hook", side_effect=invoke_hook))
stack.enter_context(patch.object(
adapter, "_create_agent", side_effect=lambda **kw: FakeAgent(kw.get("stream_delta_callback"))))
stack.enter_context(patch.object(adapter, "_finish_turn_result", side_effect=finish))
return app
def _sse_payloads(body):
return [json.loads(line[6:]) for line in body.splitlines()
if line.startswith("data: ") and line != "data: [DONE]"]
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", sorted(TRANSFORMS))
async def test_streaming_routes_deliver_transformed_final(kind):
adapter = APIServerAdapter(PlatformConfig(enabled=True, extra={}))
expected = TRANSFORMS[kind](RAW)
with ExitStack() as stack:
app = _app_with_fake_turn(adapter, TRANSFORMS[kind], stack)
async with TestClient(TestServer(app)) as cli:
r = await cli.post("/v1/chat/completions", json={
"model": "m", "messages": [{"role": "user", "content": "hi"}], "stream": True})
chunks = _sse_payloads(await r.text())
streamed = "".join(c["choices"][0]["delta"].get("content") or "" for c in chunks if c.get("choices"))
if kind == "append":
assert streamed == expected
else:
assert streamed.startswith(RAW) and streamed.endswith(expected)
r = await cli.post("/v1/responses", json={"model": "m", "input": "hi", "stream": True, "store": False})
events = _sse_payloads(await r.text())
done = [e["text"] for e in events if e.get("type") == "response.output_text.done"]
completed = [e["response"] for e in events if e.get("type") == "response.completed"]
assert done == [expected]
assert completed[0]["output"][-1]["content"][0]["text"] == expected
if kind == "append":
deltas = "".join(e["delta"] for e in events if e.get("type") == "response.output_text.delta")
assert deltas == expected