fix(api-server): deliver transform_llm_output rewrites on streaming routes
This commit is contained in:
committed by
brooklyn!
parent
7fa45eb349
commit
e7b2db974c
@@ -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"] = {
|
||||
|
||||
86
tests/gateway/test_api_server_transform_llm_output.py
Normal file
86
tests/gateway/test_api_server_transform_llm_output.py
Normal 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
|
||||
Reference in New Issue
Block a user