mirror of
https://github.com/usestrix/strix.git
synced 2026-10-10 03:28:11 +00:00
Merge d1eda1a70b into 55bc07991a
This commit is contained in:
commit
8322b4df96
4 changed files with 268 additions and 1 deletions
|
|
@ -309,6 +309,9 @@ ignore = [
|
||||||
# LiteLLM is imported lazily: the request log is wired at startup on every
|
# LiteLLM is imported lazily: the request log is wired at startup on every
|
||||||
# route, including the native OpenAI ones that never load LiteLLM.
|
# route, including the native OpenAI ones that never load LiteLLM.
|
||||||
"strix/llm/request_log.py" = ["PLC0415"]
|
"strix/llm/request_log.py" = ["PLC0415"]
|
||||||
|
# The agents SDK's chat-completions classes are imported lazily so CLI paths
|
||||||
|
# that never touch a model do not pay for them.
|
||||||
|
"strix/llm/chat_content_compat.py" = ["PLC0415"]
|
||||||
# Heavy inference deps (httpx, openai) imported lazily so auth-status checks
|
# Heavy inference deps (httpx, openai) imported lazily so auth-status checks
|
||||||
# don't pull them in.
|
# don't pull them in.
|
||||||
"strix/config/codex.py" = ["PLC0415"]
|
"strix/config/codex.py" = ["PLC0415"]
|
||||||
|
|
|
||||||
|
|
@ -42,7 +42,7 @@ from strix.config import codex
|
||||||
from strix.config.loader import load_settings
|
from strix.config.loader import load_settings
|
||||||
from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input
|
from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input
|
||||||
from strix.config.tool_call_limits import TurnToolCallLimiter
|
from strix.config.tool_call_limits import TurnToolCallLimiter
|
||||||
from strix.llm import request_log
|
from strix.llm import chat_content_compat, request_log
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
|
@ -638,6 +638,7 @@ def configure_sdk_model_defaults(settings: Settings) -> None:
|
||||||
llm = settings.llm
|
llm = settings.llm
|
||||||
set_tracing_disabled(True)
|
set_tracing_disabled(True)
|
||||||
request_log.install()
|
request_log.install()
|
||||||
|
chat_content_compat.install()
|
||||||
if codex.subscription_model(llm.model):
|
if codex.subscription_model(llm.model):
|
||||||
return
|
return
|
||||||
_configure_litellm_compatibility()
|
_configure_litellm_compatibility()
|
||||||
|
|
|
||||||
97
strix/llm/chat_content_compat.py
Normal file
97
strix/llm/chat_content_compat.py
Normal file
|
|
@ -0,0 +1,97 @@
|
||||||
|
"""Flatten list-shaped chat-completions content from OpenAI-compatible providers.
|
||||||
|
|
||||||
|
Reasoning models behind OpenAI-compatible endpoints sometimes return the
|
||||||
|
assistant content as a list of content blocks (for example
|
||||||
|
``[{"type": "thinking", ...}, {"type": "text", "text": "OK"}]``) instead of a
|
||||||
|
plain string.
|
||||||
|
|
||||||
|
``install()`` wraps the two SDK entry points that consume chat-completions
|
||||||
|
content, flattening list-shaped ``message.content`` and ``delta.content`` into
|
||||||
|
the concatenated ``text`` blocks before the SDK sees them. Non-text blocks
|
||||||
|
(e.g. ``thinking``) and unknown block types are ignored.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING, Any, cast
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
|
|
||||||
|
def _flatten_list_content(content: object) -> Any:
|
||||||
|
"""Concatenate the ``text`` blocks of list-shaped content; pass anything else through."""
|
||||||
|
if not isinstance(content, list):
|
||||||
|
return content
|
||||||
|
texts: list[str] = []
|
||||||
|
blocks = cast("list[Any]", content) # type: ignore[redundant-cast]
|
||||||
|
for part in blocks:
|
||||||
|
if not isinstance(part, dict):
|
||||||
|
continue
|
||||||
|
block = cast("dict[str, Any]", part)
|
||||||
|
if block.get("type") == "text":
|
||||||
|
text = block.get("text")
|
||||||
|
if isinstance(text, str):
|
||||||
|
texts.append(text)
|
||||||
|
return "".join(texts)
|
||||||
|
|
||||||
|
|
||||||
|
async def _flatten_chunk_stream(stream: AsyncIterator[Any]) -> AsyncIterator[Any]:
|
||||||
|
async for chunk in stream:
|
||||||
|
choices: list[Any] = getattr(chunk, "choices", None) or []
|
||||||
|
for choice in choices:
|
||||||
|
delta: Any = getattr(choice, "delta", None)
|
||||||
|
content = getattr(delta, "content", None)
|
||||||
|
if delta is not None and isinstance(content, list):
|
||||||
|
blocks = cast("list[Any]", content) # type: ignore[redundant-cast]
|
||||||
|
choice.delta = delta.model_copy(update={"content": _flatten_list_content(blocks)})
|
||||||
|
yield chunk
|
||||||
|
|
||||||
|
|
||||||
|
_installed = False
|
||||||
|
|
||||||
|
|
||||||
|
def install() -> None:
|
||||||
|
"""Wrap the SDK chat-completions content entry points (idempotent)."""
|
||||||
|
global _installed # noqa: PLW0603
|
||||||
|
if _installed:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
_install_wrappers()
|
||||||
|
except Exception: # noqa: BLE001 - a transient failure retries on the next install() call
|
||||||
|
logger.warning("could not wrap the SDK chat-completions content handling", exc_info=True)
|
||||||
|
return
|
||||||
|
_installed = True
|
||||||
|
|
||||||
|
|
||||||
|
def _install_wrappers() -> None:
|
||||||
|
from agents.models import chatcmpl_converter
|
||||||
|
from agents.models.chatcmpl_stream_handler import ChatCmplStreamHandler
|
||||||
|
|
||||||
|
# 0.19.0 names it Converter; later 0.19.x releases renamed it ChatCmplConverter.
|
||||||
|
converter_cls = getattr(chatcmpl_converter, "ChatCmplConverter", None) or (
|
||||||
|
chatcmpl_converter.Converter
|
||||||
|
)
|
||||||
|
|
||||||
|
convert = converter_cls.__dict__["message_to_output_items"].__func__
|
||||||
|
handle_stream = ChatCmplStreamHandler.__dict__["handle_stream"].__func__
|
||||||
|
|
||||||
|
def message_to_output_items(cls: Any, message: Any, *args: Any, **kwargs: Any) -> Any:
|
||||||
|
content = getattr(message, "content", None)
|
||||||
|
if isinstance(content, list):
|
||||||
|
blocks = cast("list[Any]", content) # type: ignore[redundant-cast]
|
||||||
|
message = message.model_copy(update={"content": _flatten_list_content(blocks)})
|
||||||
|
return convert(cls, message, *args, **kwargs)
|
||||||
|
|
||||||
|
def handle_list_content_stream(
|
||||||
|
cls: Any, response: Any, stream: AsyncIterator[Any], *args: Any, **kwargs: Any
|
||||||
|
) -> Any:
|
||||||
|
return handle_stream(cls, response, _flatten_chunk_stream(stream), *args, **kwargs)
|
||||||
|
|
||||||
|
converter_cls.message_to_output_items = classmethod(message_to_output_items) # type: ignore[method-assign]
|
||||||
|
ChatCmplStreamHandler.handle_stream = classmethod(handle_list_content_stream) # type: ignore[assignment]
|
||||||
|
_installed = True
|
||||||
166
tests/test_list_content_compat.py
Normal file
166
tests/test_list_content_compat.py
Normal file
|
|
@ -0,0 +1,166 @@
|
||||||
|
"""Tests for list-shaped chat-completions content from OpenAI-compatible gateways.
|
||||||
|
|
||||||
|
Reasoning models behind some OpenAI-compatible endpoints return assistant
|
||||||
|
content as a list of content blocks (``thinking`` + ``text``) instead of a
|
||||||
|
plain string. The OpenAI client constructs its models without validating that
|
||||||
|
field, and the SDK then crashes on ``ResponseOutputText``/``ResponseTextDeltaEvent``
|
||||||
|
validation. ``chat_content_compat.install()`` flattens the text blocks before
|
||||||
|
the SDK consumes them; a local gateway that answers with list-shaped content
|
||||||
|
proves both the non-streaming and the streaming path work end to end.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import threading
|
||||||
|
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from agents.model_settings import ModelSettings
|
||||||
|
from agents.models.interface import ModelTracing
|
||||||
|
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
|
||||||
|
from openai import AsyncOpenAI
|
||||||
|
from openai.types.responses import (
|
||||||
|
ResponseCompletedEvent,
|
||||||
|
ResponseOutputMessage,
|
||||||
|
ResponseOutputText,
|
||||||
|
ResponseTextDeltaEvent,
|
||||||
|
)
|
||||||
|
|
||||||
|
from strix.llm.chat_content_compat import install as install_content_compat
|
||||||
|
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import AsyncIterator, Iterator
|
||||||
|
|
||||||
|
install_content_compat()
|
||||||
|
|
||||||
|
|
||||||
|
def _list_content_completion() -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": "chatcmpl-1",
|
||||||
|
"object": "chat.completion",
|
||||||
|
"created": 0,
|
||||||
|
"model": "gw-model",
|
||||||
|
"choices": [
|
||||||
|
{
|
||||||
|
"index": 0,
|
||||||
|
"finish_reason": "stop",
|
||||||
|
"message": {
|
||||||
|
"role": "assistant",
|
||||||
|
"content": [
|
||||||
|
{"type": "thinking", "thinking": "considering the request"},
|
||||||
|
{"type": "text", "text": "OK"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _list_content_chunks() -> list[dict[str, Any]]:
|
||||||
|
def chunk(delta: dict[str, Any], finish_reason: str | None = None) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"id": "chatcmpl-2",
|
||||||
|
"object": "chat.completion.chunk",
|
||||||
|
"created": 0,
|
||||||
|
"model": "gw-model",
|
||||||
|
"choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}],
|
||||||
|
}
|
||||||
|
|
||||||
|
return [
|
||||||
|
chunk({"role": "assistant"}),
|
||||||
|
chunk({"content": [{"type": "thinking", "thinking": "hmm"}]}),
|
||||||
|
chunk({"content": [{"type": "text", "text": "OK"}]}),
|
||||||
|
chunk({}, "stop"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class _Handler(BaseHTTPRequestHandler):
|
||||||
|
"""A gateway whose assistant content arrives as a list of content blocks."""
|
||||||
|
|
||||||
|
def log_message(self, *args: Any) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def do_POST(self) -> None:
|
||||||
|
length = int(self.headers.get("Content-Length", 0))
|
||||||
|
body = json.loads(self.rfile.read(length) or b"{}")
|
||||||
|
if body.get("stream"):
|
||||||
|
payload = "".join(f"data: {json.dumps(c)}\n\n" for c in _list_content_chunks())
|
||||||
|
payload += "data"
|
||||||
|
payload_bytes = payload.encode()
|
||||||
|
self.send_response(200)
|
||||||
|
self.send_header("Content-Type", "text/event-stream")
|
||||||
|
else:
|
||||||
|
payload_bytes = json.dumps(_list_content_completion()).encode()
|
||||||
|
self.send_response(200)
|
||||||
|
self.send_header("Content-Type", "application/json")
|
||||||
|
self.send_header("Content-Length", str(len(payload_bytes)))
|
||||||
|
self.end_headers()
|
||||||
|
self.wfile.write(payload_bytes)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def gateway_url() -> Iterator[str]:
|
||||||
|
server = HTTPServer(("127.0.0.1", 0), _Handler)
|
||||||
|
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||||
|
thread.start()
|
||||||
|
try:
|
||||||
|
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
|
||||||
|
finally:
|
||||||
|
server.shutdown()
|
||||||
|
server.server_close()
|
||||||
|
|
||||||
|
|
||||||
|
def _model(base_url: str) -> OpenAIChatCompletionsModel:
|
||||||
|
client = AsyncOpenAI(api_key="tok", base_url=base_url)
|
||||||
|
return OpenAIChatCompletionsModel(model="gw-model", openai_client=client)
|
||||||
|
|
||||||
|
|
||||||
|
def _call_kwargs() -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"system_instructions": "s",
|
||||||
|
"input": "hi",
|
||||||
|
"model_settings": ModelSettings(),
|
||||||
|
"tools": [],
|
||||||
|
"output_schema": None,
|
||||||
|
"handoffs": [],
|
||||||
|
"tracing": ModelTracing.DISABLED,
|
||||||
|
"previous_response_id": None,
|
||||||
|
"conversation_id": None,
|
||||||
|
"prompt": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _drain(gen: AsyncIterator[Any]) -> list[Any]:
|
||||||
|
return [event async for event in gen]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_get_response_flattens_list_content(gateway_url: str) -> None:
|
||||||
|
response = await _model(gateway_url).get_response(**_call_kwargs())
|
||||||
|
|
||||||
|
message = response.output[0]
|
||||||
|
assert isinstance(message, ResponseOutputMessage)
|
||||||
|
text = message.content[0]
|
||||||
|
assert isinstance(text, ResponseOutputText)
|
||||||
|
assert text.text == "OK"
|
||||||
|
|
||||||
|
assert response.usage is not None
|
||||||
|
assert response.usage.total_tokens == 8
|
||||||
|
|
||||||
|
|
||||||
|
async def test_stream_response_flattens_list_content_deltas(gateway_url: str) -> None:
|
||||||
|
events = await _drain(_model(gateway_url).stream_response(**_call_kwargs()))
|
||||||
|
|
||||||
|
deltas = [e for e in events if isinstance(e, ResponseTextDeltaEvent)]
|
||||||
|
assert [d.delta for d in deltas] == ["OK"]
|
||||||
|
|
||||||
|
completed = events[-1]
|
||||||
|
assert isinstance(completed, ResponseCompletedEvent)
|
||||||
|
message = completed.response.output[0]
|
||||||
|
assert isinstance(message, ResponseOutputMessage)
|
||||||
|
text = message.content[0]
|
||||||
|
assert isinstance(text, ResponseOutputText)
|
||||||
|
assert text.text == "OK"
|
||||||
Loading…
Add table
Reference in a new issue