From 980216860e4965928a992cc6114a038587dd9291 Mon Sep 17 00:00:00 2001 From: Ahmed Allam Date: Thu, 30 Jul 2026 05:04:19 +0000 Subject: [PATCH] feat(llm): opt-in LLM_DISABLE_STREAMING for non-streaming OpenAI-compatible endpoints Some OpenAI-compatible gateways don't support Server-Sent Events (or deliver them unreliably), but the SDK run loop Strix uses only issues streamed requests, so such a gateway fails every turn. Add an opt-in LLM_DISABLE_STREAMING setting that wraps the resolved model in _NonStreamingModel: each turn makes one non-streaming get_response and replays the completed result as a single terminal stream event, so tool calls, usage, and the rest of the agent loop are unchanged. Subscription (ChatGPT) models are always streamed and are not wrapped. --- README.md | 1 + pyproject.toml | 1 + strix/config/models.py | 143 +++++++++++++++++- strix/config/settings.py | 4 + tests/test_disable_streaming.py | 258 ++++++++++++++++++++++++++++++++ 5 files changed, 404 insertions(+), 3 deletions(-) create mode 100644 tests/test_disable_streaming.py diff --git a/README.md b/README.md index 982a635d..47fad3b2 100644 --- a/README.md +++ b/README.md @@ -262,6 +262,7 @@ export LLM_API_KEY="your-api-key" export LLM_API_BASE="your-api-base-url" # if using a local model, e.g. Ollama, LMStudio export PERPLEXITY_API_KEY="your-api-key" # for search capabilities export STRIX_REASONING_EFFORT="high" # control thinking effort (default: high, quick scan: medium) +export LLM_DISABLE_STREAMING="true" # for OpenAI-compatible endpoints that don't support streaming ``` > [!NOTE] diff --git a/pyproject.toml b/pyproject.toml index c9ea8bc1..a24b6a12 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -220,6 +220,7 @@ ignore = [ # Stdlib HTTP handler overrides (do_GET/do_POST). "strix/interface/auth_cli.py" = ["N802"] "tests/test_codex_streaming.py" = ["N802"] +"tests/test_disable_streaming.py" = ["N802"] "tests/test_report_pdf.py" = ["S105", "S106"] # Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a # circular dependency with strix.telemetry / strix.interface.viewer.report_pdf. diff --git a/strix/config/models.py b/strix/config/models.py index 7892957d..1c84ef3c 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -5,6 +5,7 @@ from __future__ import annotations import contextlib import inspect import os +import time from typing import TYPE_CHECKING, Any from agents import ( @@ -13,6 +14,8 @@ from agents import ( set_tracing_disabled, ) from agents.model_settings import ModelSettings +from agents.models.fake_id import FAKE_RESPONSES_ID +from agents.models.interface import Model from agents.models.multi_provider import MultiProvider from agents.models.openai_responses import OpenAIResponsesModel from agents.retry import ( @@ -21,6 +24,8 @@ from agents.retry import ( RetryPolicyContext, retry_policies, ) +from openai.types.responses import Response, ResponseCompletedEvent +from openai.types.responses.response_usage import ResponseUsage from openai.types.shared import Reasoning from strix.config import codex @@ -30,8 +35,15 @@ from strix.config.loader import load_settings if TYPE_CHECKING: from collections.abc import AsyncIterator - from agents.models.interface import Model, ModelProvider + from agents.agent_output import AgentOutputSchemaBase + from agents.handoffs import Handoff + from agents.items import ModelResponse, TResponseInputItem, TResponseStreamEvent + from agents.models.interface import ModelProvider, ModelTracing + from agents.retry import ModelRetryAdvice, ModelRetryAdviceRequest + from agents.tool import Tool + from agents.usage import Usage from openai import AsyncOpenAI + from openai.types.responses.response_prompt_param import ResponsePromptParam from strix.config.settings import LlmSettings, ReasoningEffort, Settings @@ -135,6 +147,124 @@ class _CodexResponsesModel(OpenAIResponsesModel): await result +class _NonStreamingModel(Model): + """Serve the SDK's streamed run loop from a single non-streaming request. + + Some OpenAI-compatible gateways do not support Server-Sent Events, or + deliver them unreliably (dropping structured tool-call deltas, or stalling + mid-stream so the whole turn waits out the read timeout). The SDK run loop + Strix uses only issues streamed requests, so such a gateway fails every + turn. Opt in with ``LLM_DISABLE_STREAMING=true`` to wrap the resolved model + so each turn makes one non-streaming ``get_response`` (``stream:false`` on + the wire) and the completed result is replayed as a single terminal stream + event. The run loop then executes tools and emits run items from that final + response exactly as it would for a real stream, so nothing else changes. + """ + + def __init__(self, inner: Model) -> None: + self._inner = inner + + async def close(self) -> None: + await self._inner.close() + + def get_retry_advice(self, request: ModelRetryAdviceRequest) -> ModelRetryAdvice | None: + return self._inner.get_retry_advice(request) + + async def get_response( + self, + system_instructions: str | None, + input: str | list[TResponseInputItem], # noqa: A002 + model_settings: ModelSettings, + tools: list[Tool], + output_schema: AgentOutputSchemaBase | None, + handoffs: list[Handoff], + tracing: ModelTracing, + *, + previous_response_id: str | None, + conversation_id: str | None, + prompt: ResponsePromptParam | None, + ) -> ModelResponse: + return await self._inner.get_response( + system_instructions, + input, + model_settings, + tools, + output_schema, + handoffs, + tracing, + previous_response_id=previous_response_id, + conversation_id=conversation_id, + prompt=prompt, + ) + + async def stream_response( + self, + system_instructions: str | None, + input: str | list[TResponseInputItem], # noqa: A002 + model_settings: ModelSettings, + tools: list[Tool], + output_schema: AgentOutputSchemaBase | None, + handoffs: list[Handoff], + tracing: ModelTracing, + *, + previous_response_id: str | None, + conversation_id: str | None, + prompt: ResponsePromptParam | None, + ) -> AsyncIterator[TResponseStreamEvent]: + response = await self._inner.get_response( + system_instructions, + input, + model_settings, + tools, + output_schema, + handoffs, + tracing, + previous_response_id=previous_response_id, + conversation_id=conversation_id, + prompt=prompt, + ) + yield _completed_stream_event(response, getattr(self._inner, "model", None)) + + +def _completed_stream_event( + model_response: ModelResponse, model_name: object | None +) -> TResponseStreamEvent: + """Wrap a non-streamed ``ModelResponse`` as the terminal event of a stream. + + The run loop builds its authoritative per-turn response solely from the + ``response.completed`` event, so a single event carrying the full output + and usage is all it needs. + """ + response = Response( + id=model_response.response_id or FAKE_RESPONSES_ID, + created_at=time.time(), + model=str(model_name) if model_name else "", + object="response", + output=list(model_response.output), + tool_choice="auto", + tools=[], + parallel_tool_calls=False, + usage=_response_usage(model_response.usage), + ) + return ResponseCompletedEvent( + response=response, + sequence_number=0, + type="response.completed", + ) + + +def _response_usage(usage: Usage | None) -> ResponseUsage | None: + if usage is None: + return None + return ResponseUsage( + input_tokens=usage.input_tokens, + output_tokens=usage.output_tokens, + total_tokens=usage.total_tokens, + input_tokens_details=usage.input_tokens_details, + output_tokens_details=usage.output_tokens_details, + ) + + class StrixProvider(MultiProvider): """Route any non-OpenAI prefix through LiteLLM with the prefix preserved, so users type ``deepseek/deepseek-chat`` rather than @@ -159,14 +289,21 @@ class StrixProvider(MultiProvider): return self._get_fallback_provider("litellm"), original_model_name def get_model(self, model_name: str | None) -> Model: + llm = load_settings().llm slug = codex.subscription_model(model_name) if slug: + # The ChatGPT subscription backend is always streamed; it has no + # non-streaming mode to fall back to, so LLM_DISABLE_STREAMING + # does not apply here. return _CodexResponsesModel( slug, codex.get_subscription_client(), - reasoning_effort=load_settings().llm.reasoning_effort, + reasoning_effort=llm.reasoning_effort, ) - return super().get_model(model_name) + model = super().get_model(model_name) + if llm.disable_streaming: + return _NonStreamingModel(model) + return model DEFAULT_MODEL_RETRY = ModelRetrySettings( diff --git a/strix/config/settings.py b/strix/config/settings.py index 016a8ad9..e53d125c 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -48,6 +48,10 @@ class LlmSettings(BaseSettings): default=True, alias="STRIX_PROMPT_CACHE", ) + disable_streaming: bool = Field( + default=False, + alias="LLM_DISABLE_STREAMING", + ) timeout: int = Field(default=300, alias="LLM_TIMEOUT") diff --git a/tests/test_disable_streaming.py b/tests/test_disable_streaming.py new file mode 100644 index 00000000..16afafbf --- /dev/null +++ b/tests/test_disable_streaming.py @@ -0,0 +1,258 @@ +"""Tests for LLM_DISABLE_STREAMING: serve the streamed run loop without SSE. + +A gateway that rejects ``stream:true`` (or delivers SSE unreliably) breaks the +SDK run loop, which only issues streamed requests. ``_NonStreamingModel`` wraps +the resolved model so each turn makes one non-streaming ``get_response`` and +replays the completed result as a single terminal stream event. A local server +that rejects streamed requests but answers non-streamed ones — including a +structured tool call — proves the wrapper works where the stock model fails. +""" + +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 Model, ModelTracing +from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel +from openai import AsyncOpenAI, BadRequestError +from openai.types.responses import ( + ResponseCompletedEvent, + ResponseFunctionToolCall, + ResponseOutputMessage, + ResponseOutputText, +) + +from strix.config import codex, loader +from strix.config.loader import load_settings +from strix.config.models import StrixProvider, _NonStreamingModel + + +if TYPE_CHECKING: + from collections.abc import AsyncIterator, Iterator + + +def _tool_call_completion() -> dict[str, Any]: + return { + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 0, + "model": "gw-model", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "do_thing", "arguments": '{"n": 1}'}, + } + ], + }, + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, + } + + +def _text_completion() -> dict[str, Any]: + return { + "id": "chatcmpl-2", + "object": "chat.completion", + "created": 0, + "model": "gw-model", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "hello from gateway"}, + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + + +_CAPTURED: dict[str, Any] = {} +_PAYLOAD: dict[str, dict[str, Any]] = {"value": _tool_call_completion()} + + +class _Handler(BaseHTTPRequestHandler): + """A gateway that only speaks non-streaming Chat Completions.""" + + 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"{}") + _CAPTURED.clear() + _CAPTURED.update(body) + if body.get("stream"): + payload = json.dumps( + {"error": {"message": "streaming is not supported by this endpoint"}} + ).encode() + self.send_response(400) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(payload))) + self.end_headers() + self.wfile.write(payload) + return + payload = json.dumps(_PAYLOAD["value"]).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(payload))) + self.end_headers() + self.wfile.write(payload) + + +@pytest.fixture +def gateway_url() -> Iterator[str]: + _PAYLOAD["value"] = _tool_call_completion() + 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] + + +@pytest.mark.asyncio +async def test_stock_model_streaming_fails_on_non_streaming_gateway(gateway_url: str) -> None: + # The stock model issues stream:true and the gateway rejects it. + model = _model(gateway_url) + with pytest.raises(BadRequestError, match="streaming is not supported"): + await _drain(model.stream_response(**_call_kwargs())) + assert _CAPTURED["stream"] is True + + +@pytest.mark.asyncio +async def test_wrapper_streams_tool_call_without_streaming_request(gateway_url: str) -> None: + # The wrapper turns the streamed run-loop call into one non-streaming + # request and replays the completed result as a terminal stream event. + model = _NonStreamingModel(_model(gateway_url)) + events = await _drain(model.stream_response(**_call_kwargs())) + + assert _CAPTURED.get("stream") is not True + assert len(events) == 1 + completed = events[0] + assert isinstance(completed, ResponseCompletedEvent) + + tool_call = completed.response.output[0] + assert isinstance(tool_call, ResponseFunctionToolCall) + assert tool_call.name == "do_thing" + assert json.loads(tool_call.arguments) == {"n": 1} + + assert completed.response.usage is not None + assert completed.response.usage.total_tokens == 7 + + +@pytest.mark.asyncio +async def test_wrapper_streams_plain_text(gateway_url: str) -> None: + _PAYLOAD["value"] = _text_completion() + model = _NonStreamingModel(_model(gateway_url)) + events = await _drain(model.stream_response(**_call_kwargs())) + + assert _CAPTURED.get("stream") is not True + message = events[0].response.output[0] + assert isinstance(message, ResponseOutputMessage) + text = message.content[0] + assert isinstance(text, ResponseOutputText) + assert text.text == "hello from gateway" + + +@pytest.mark.asyncio +async def test_wrapper_get_response_stays_non_streaming(gateway_url: str) -> None: + # The non-streaming path is a plain pass-through to the inner model. + model = _NonStreamingModel(_model(gateway_url)) + response = await model.get_response(**_call_kwargs()) + assert _CAPTURED.get("stream") is not True + tool_call = response.output[0] + assert isinstance(tool_call, ResponseFunctionToolCall) + assert tool_call.name == "do_thing" + + +class _DummyModel(Model): + async def get_response(self, *args: Any, **kwargs: Any) -> Any: + raise NotImplementedError + + def stream_response(self, *args: Any, **kwargs: Any) -> Any: + raise NotImplementedError + + +@pytest.fixture +def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: + for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING"): + monkeypatch.delenv(key, raising=False) + monkeypatch.setattr(loader, "_cached", None) + monkeypatch.setattr(loader, "_override", None) + yield + + +def test_get_model_wraps_when_disabled( + monkeypatch: pytest.MonkeyPatch, _reset_settings: None +) -> None: + inner = _DummyModel() + monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: inner) + monkeypatch.setenv("LLM_DISABLE_STREAMING", "true") + load_settings() + + model = StrixProvider().get_model("openai/gpt-4o-mini") + assert isinstance(model, _NonStreamingModel) + + +def test_get_model_unwrapped_by_default( + monkeypatch: pytest.MonkeyPatch, _reset_settings: None +) -> None: + inner = _DummyModel() + monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: inner) + load_settings() + + model = StrixProvider().get_model("openai/gpt-4o-mini") + assert model is inner + + +def test_get_model_does_not_wrap_subscription_model( + monkeypatch: pytest.MonkeyPatch, _reset_settings: None +) -> None: + # Subscription (ChatGPT) models are always streamed and must not be wrapped. + monkeypatch.setattr(codex, "subscription_model", lambda *_: "gpt-5.5") + monkeypatch.setattr(codex, "get_subscription_client", lambda: AsyncOpenAI(api_key="x")) + monkeypatch.setenv("LLM_DISABLE_STREAMING", "true") + load_settings() + + model = StrixProvider().get_model("gpt-5.5") + assert not isinstance(model, _NonStreamingModel)