Compare commits

...
Author SHA1 Message Date
Alex Schapiro d8ad8e3572 feat(models): STRIX_STREAM_MODE opt-in non-streaming for custom endpoints
Replace the api_base-based non-streaming heuristic with an explicit
STRIX_STREAM_MODE=auto|always|never setting. auto/always stream (unchanged
default); never routes through the non-streaming wrapper for endpoints whose
streamed responses drop tool calls.
2026-07-30 00:51:43 +00:00
Alex Schapiro 788a5393db docs(models): generalize non-streaming wrapper docstring 2026-07-29 15:18:26 +00:00
Alex Schapiro 7b82ff8432 fix(models): run custom OpenAI-compatible endpoints non-streamed
Some OpenAI-compatible endpoints return valid tool_calls for a non-streamed
completion but, when streamed, emit the tool call as plain text or drop it and
close the stream, leaving Strix's tool-driven loop with nothing to execute.

Wrap the model for custom (api_base) endpoints so the request is made
non-streamed (where tool calling works) while still presenting the streaming
interface the runner consumes. Hosted providers are unchanged. Opt back into
streaming with STRIX_STREAM_CUSTOM_ENDPOINT=1.
2026-07-29 14:45:37 +00:00
3 changed files with 295 additions and 3 deletions
+115 -3
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import contextlib import contextlib
import inspect import inspect
import os import os
import time
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from agents import ( from agents import (
@@ -13,6 +14,8 @@ from agents import (
set_tracing_disabled, set_tracing_disabled,
) )
from agents.model_settings import ModelSettings 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.multi_provider import MultiProvider
from agents.models.openai_responses import OpenAIResponsesModel from agents.models.openai_responses import OpenAIResponsesModel
from agents.retry import ( from agents.retry import (
@@ -21,6 +24,12 @@ from agents.retry import (
RetryPolicyContext, RetryPolicyContext,
retry_policies, retry_policies,
) )
from openai.types.responses import (
Response,
ResponseCompletedEvent,
ResponseOutputItemDoneEvent,
ResponseUsage,
)
from openai.types.shared import Reasoning from openai.types.shared import Reasoning
from strix.config import codex from strix.config import codex
@@ -30,8 +39,14 @@ from strix.config.loader import load_settings
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncIterator 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.tool import Tool
from agents.usage import Usage
from openai import AsyncOpenAI from openai import AsyncOpenAI
from openai.types.responses import ResponsePromptParam
from strix.config.settings import ReasoningEffort, Settings from strix.config.settings import ReasoningEffort, Settings
@@ -135,6 +150,99 @@ class _CodexResponsesModel(OpenAIResponsesModel):
await result await result
def _to_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 _NonStreamingModel(Model):
"""Run a model non-streamed but expose the streaming interface the runner uses.
Some OpenAI-compatible endpoints (notably gateways serving reasoning models)
return valid ``tool_calls`` for a non-streamed completion but, when streamed,
emit the tool call as plain text or drop it and close the stream — leaving
Strix's tool-driven loop with nothing to execute. Selecting
``STRIX_STREAM_MODE=never`` routes through this wrapper, which makes the real
request non-streamed (where tool calling works) and synthesizes the minimal
event sequence the runner consumes from a stream, so the rest of the pipeline
is unchanged. The only user-visible difference is no token-by-token output.
"""
def __init__(self, inner: Model) -> None:
self._inner = inner
async def get_response(self, *args: Any, **kwargs: Any) -> ModelResponse:
return await self._inner.get_response(*args, **kwargs)
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 = None,
conversation_id: str | None = None,
prompt: ResponsePromptParam | None = None,
) -> AsyncIterator[TResponseStreamEvent]:
model_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,
)
sequence = 0
for index, item in enumerate(model_response.output):
yield ResponseOutputItemDoneEvent(
item=item,
output_index=index,
type="response.output_item.done",
sequence_number=sequence,
)
sequence += 1
response = Response(
id=model_response.response_id or FAKE_RESPONSES_ID,
created_at=time.time(),
model=str(getattr(self._inner, "model", "")),
object="response",
output=model_response.output,
tool_choice="auto",
tools=[],
parallel_tool_calls=False,
usage=_to_response_usage(model_response.usage),
)
yield ResponseCompletedEvent(
response=response,
type="response.completed",
sequence_number=sequence,
)
def get_retry_advice(self, request: Any) -> Any:
return self._inner.get_retry_advice(request)
def _should_run_non_streamed(settings: Settings) -> bool:
return settings.llm.stream_mode == "never"
class StrixProvider(MultiProvider): class StrixProvider(MultiProvider):
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved, """Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
so users type ``deepseek/deepseek-chat`` rather than so users type ``deepseek/deepseek-chat`` rather than
@@ -159,14 +267,18 @@ class StrixProvider(MultiProvider):
return self._get_fallback_provider("litellm"), original_model_name return self._get_fallback_provider("litellm"), original_model_name
def get_model(self, model_name: str | None) -> Model: def get_model(self, model_name: str | None) -> Model:
settings = load_settings()
slug = codex.subscription_model(model_name) slug = codex.subscription_model(model_name)
if slug: if slug:
return _CodexResponsesModel( return _CodexResponsesModel(
slug, slug,
codex.get_subscription_client(), codex.get_subscription_client(),
reasoning_effort=load_settings().llm.reasoning_effort, reasoning_effort=settings.llm.reasoning_effort,
) )
return super().get_model(model_name) model = super().get_model(model_name)
if _should_run_non_streamed(settings):
return _NonStreamingModel(model)
return model
DEFAULT_MODEL_RETRY = ModelRetrySettings( DEFAULT_MODEL_RETRY = ModelRetrySettings(
+7
View File
@@ -9,6 +9,7 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
ReasoningEffort = Literal["none", "minimal", "low", "medium", "high", "xhigh"] ReasoningEffort = Literal["none", "minimal", "low", "medium", "high", "xhigh"]
StreamMode = Literal["auto", "always", "never"]
_BASE_CONFIG = SettingsConfigDict( _BASE_CONFIG = SettingsConfigDict(
case_sensitive=False, case_sensitive=False,
@@ -40,6 +41,12 @@ class LlmSettings(BaseSettings):
default=False, default=False,
alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE", alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE",
) )
# auto/always stream; never runs non-streamed, for endpoints that stream tool
# calls incorrectly (some OpenAI-compatible gateways serving reasoning models).
stream_mode: StreamMode = Field(
default="auto",
alias="STRIX_STREAM_MODE",
)
prompt_cache: bool = Field( prompt_cache: bool = Field(
default=True, default=True,
alias="STRIX_PROMPT_CACHE", alias="STRIX_PROMPT_CACHE",
+173
View File
@@ -0,0 +1,173 @@
"""Tests for the non-streaming wrapper used on custom OpenAI-compatible endpoints."""
from __future__ import annotations
from types import SimpleNamespace
from typing import TYPE_CHECKING, cast
from unittest.mock import patch
import pytest
from agents.items import ModelResponse
from agents.model_settings import ModelSettings
from agents.models.interface import Model, ModelTracing
from agents.usage import Usage
from openai.types.responses import (
ResponseCompletedEvent,
ResponseFunctionToolCall,
ResponseOutputItemDoneEvent,
ResponseStreamEvent,
)
from strix.config.models import StrixProvider, _NonStreamingModel, _to_response_usage
if TYPE_CHECKING:
from collections.abc import AsyncIterator
def _tool_call() -> ResponseFunctionToolCall:
return ResponseFunctionToolCall(
arguments='{"command": "ls"}',
call_id="call_1",
name="terminal_execute",
type="function_call",
)
class _FakeModel(Model):
def __init__(self, response: ModelResponse) -> None:
self.model = "fake-model"
self._response = response
self.get_response_calls = 0
async def get_response(self, *_args: object, **_kwargs: object) -> ModelResponse:
self.get_response_calls += 1
return self._response
async def stream_response( # pragma: no cover
self, *_args: object, **_kwargs: object
) -> AsyncIterator[ResponseStreamEvent]:
for _ in range(0):
yield cast("ResponseStreamEvent", None)
raise AssertionError("inner stream_response must never be called")
def _settings(*, api_base: str | None, stream_mode: str = "auto") -> SimpleNamespace:
return SimpleNamespace(
llm=SimpleNamespace(
model="openai/glm",
api_base=api_base,
stream_mode=stream_mode,
reasoning_effort="high",
)
)
def test_to_response_usage_maps_token_details() -> None:
usage = Usage(requests=1, input_tokens=10, output_tokens=5, total_tokens=15)
usage.input_tokens_details.cached_tokens = 4
usage.output_tokens_details.reasoning_tokens = 3
mapped = _to_response_usage(usage)
assert mapped is not None
assert (mapped.input_tokens, mapped.output_tokens, mapped.total_tokens) == (10, 5, 15)
assert mapped.input_tokens_details.cached_tokens == 4
assert mapped.output_tokens_details.reasoning_tokens == 3
def test_to_response_usage_none() -> None:
assert _to_response_usage(None) is None
@pytest.mark.asyncio
async def test_stream_response_synthesizes_tool_call_from_non_streamed() -> None:
tool_call = _tool_call()
inner = _FakeModel(
ModelResponse(
output=[tool_call],
usage=Usage(requests=1, input_tokens=10, output_tokens=5, total_tokens=15),
response_id="resp_123",
)
)
wrapper = _NonStreamingModel(inner)
events = [
event
async for event in wrapper.stream_response(
"sys",
"hi",
ModelSettings(),
[],
None,
[],
ModelTracing.DISABLED,
)
]
assert inner.get_response_calls == 1
item_done = [e for e in events if isinstance(e, ResponseOutputItemDoneEvent)]
completed = [e for e in events if isinstance(e, ResponseCompletedEvent)]
assert len(item_done) == 1
assert item_done[0].item == tool_call
assert len(completed) == 1
final = completed[0].response
assert final.output == [tool_call]
assert final.id == "resp_123"
assert final.usage is not None
assert final.usage.total_tokens == 15
# sequence numbers are strictly increasing
assert [e.sequence_number for e in events] == list(range(len(events)))
@pytest.mark.asyncio
async def test_stream_response_delegates_get_response() -> None:
inner = _FakeModel(
ModelResponse(output=[], usage=Usage(), response_id=None),
)
wrapper = _NonStreamingModel(inner)
result = await wrapper.get_response(
"sys", "hi", ModelSettings(), [], None, [], ModelTracing.DISABLED
)
assert result is inner._response
assert inner.get_response_calls == 1
def test_get_model_auto_streams_custom_endpoint() -> None:
sentinel = _FakeModel(ModelResponse(output=[], usage=Usage(), response_id=None))
with (
patch("strix.config.models.load_settings", return_value=_settings(api_base="http://x/v1")),
patch(
"agents.models.multi_provider.MultiProvider.get_model",
return_value=sentinel,
),
):
model = StrixProvider().get_model("openai/glm")
assert model is sentinel
def test_get_model_auto_streams_hosted() -> None:
sentinel = _FakeModel(ModelResponse(output=[], usage=Usage(), response_id=None))
with (
patch("strix.config.models.load_settings", return_value=_settings(api_base="")),
patch(
"agents.models.multi_provider.MultiProvider.get_model",
return_value=sentinel,
),
):
model = StrixProvider().get_model("openai/gpt-4o")
assert model is sentinel
def test_get_model_stream_mode_never_wraps() -> None:
sentinel = _FakeModel(ModelResponse(output=[], usage=Usage(), response_id=None))
with (
patch(
"strix.config.models.load_settings",
return_value=_settings(api_base="http://x/v1", stream_mode="never"),
),
patch(
"agents.models.multi_provider.MultiProvider.get_model",
return_value=sentinel,
),
):
model = StrixProvider().get_model("openai/glm")
assert isinstance(model, _NonStreamingModel)