Compare commits

...
Author SHA1 Message Date
Alex Schapiro 0e28732d29 feat(warmup): fail fast when a streamed request yields no tool call
OpenAI-compatible endpoints that return valid tool_calls for a
non-streamed request but emit plain text (or drop the call) when
streamed leave Strix's tool-driven scan unable to act. Probe the
streaming path up front for custom endpoints / Ollama and abort with
actionable guidance; STRIX_SKIP_TOOL_CALL_PROBE opts out.
2026-07-29 14:17:17 +00:00
alex sandGitHub 9de747d135 fix(cost): capture OpenRouter streamed usage.cost (fixes $0 kimi-k3 c… (#929)
* fix(cost): capture OpenRouter streamed usage.cost (fixes $0 kimi-k3 cost)

* refactor(cost): encapsulate streamed OpenRouter cost cache, clear per run

* test(cost): resolve OpenRouter handler via LiteLLM provider pipeline
2026-07-28 23:28:34 -04:00
7 changed files with 598 additions and 2 deletions
+45
View File
@@ -277,6 +277,51 @@ def _configure_litellm_compatibility() -> None:
litellm.suppress_debug_info = True
_register_litellm_cost_callback()
_install_openrouter_stream_cost_capture()
def _install_openrouter_stream_cost_capture() -> None:
"""Preserve OpenRouter's per-stream cost, which LiteLLM drops when streaming.
OpenRouter reports the real charge in ``usage.cost`` of the final stream
chunk, but LiteLLM rebuilds streamed responses from token-only fields and
discards it (its non-streamed path stashes the cost in hidden params; the
streaming path does not). Every scan streams, so without this the cost is
lost and Strix falls back to a cost-map estimate that is missing entirely
for new models (e.g. kimi-k3), reporting $0. Subclass the OpenRouter
streaming handler to record the cost keyed by response id so the cost
callback can recover the exact charge for the matching rebuilt response.
"""
import litellm
from litellm.llms.openrouter.chat.transformation import (
OpenRouterChatCompletionStreamingHandler,
OpenrouterConfig,
)
from strix.report.state import streamed_openrouter_costs
class _StrixOpenRouterStreamingHandler(OpenRouterChatCompletionStreamingHandler):
def chunk_parser(self, chunk: dict[str, Any]) -> Any:
stream = super().chunk_parser(chunk)
streamed_openrouter_costs.remember(
chunk.get("id") or getattr(stream, "id", None), chunk.get("usage")
)
return stream
class _StrixOpenrouterConfig(OpenrouterConfig):
def get_model_response_iterator(
self, streaming_response: Any, sync_stream: bool, json_mode: bool | None = False
) -> Any:
return _StrixOpenRouterStreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
# LiteLLM's provider-config factory reads litellm.OpenrouterConfig at call
# time, so overriding the attribute is enough for the subclass to take
# effect. (type: ignore — mypy rejects reassigning a class attribute.)
litellm.OpenrouterConfig = _StrixOpenrouterConfig # type: ignore[misc]
_OPENROUTER_ATTRIBUTION_HEADERS = {
+4
View File
@@ -40,6 +40,10 @@ class LlmSettings(BaseSettings):
default=False,
alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE",
)
skip_tool_call_probe: bool = Field(
default=False,
alias="STRIX_SKIP_TOOL_CALL_PROBE",
)
prompt_cache: bool = Field(
default=True,
alias="STRIX_PROMPT_CACHE",
+190
View File
@@ -0,0 +1,190 @@
"""Preflight probe: verify the model emits structured tool calls when streamed.
Strix is entirely tool-driven and runs every agent turn as a streamed request.
Some OpenAI-compatible endpoints return a valid ``tool_calls`` response when a
completion is requested non-streamed, but under streaming they emit the tool
call as plain assistant text (or drop it entirely) and close the stream. The
Agents SDK then sees a normal final message, the scan makes no progress, and it
either stalls waiting for input or burns turns on empty output.
There is no safe client-side way to execute a tool call the endpoint never
streamed, so we detect the missing capability up front — using the same
streaming path the scan uses — and fail loudly with actionable guidance.
"""
from __future__ import annotations
import asyncio
import logging
from typing import TYPE_CHECKING, Any
from agents import ModelSettings, ModelTracing
from agents.tool import FunctionTool
from openai.types.responses import ResponseFunctionToolCall
from strix.config.models import StrixProvider
if TYPE_CHECKING:
from strix.config.settings import Settings
logger = logging.getLogger(__name__)
_PROBE_RETRIES = 2
_GUIDANCE = (
"The configured LLM endpoint did not return a structured tool call when "
"streamed.\n\n"
"Strix drives every action through native `tool_calls`, and it streams every "
"turn. Some OpenAI-compatible servers return tool calls correctly for a "
"non-streamed request but, when streamed, emit the tool call as plain text "
"(or omit it) and end the response — so Strix can never act.\n\n"
"Fixes:\n"
" - llama.cpp / llama-server: start with `--jinja` so the chat template "
"produces streamed `tool_calls` deltas.\n"
" - Ollama: use a model whose template wires tool calling, and disable "
"'thinking' if the template can't stream tools alongside it.\n"
" - vLLM: set a matching `--tool-call-parser` (and `--enable-auto-tool-choice`) "
"for the served model.\n"
" - Other gateways: confirm streamed tool calling works for this model "
"(a non-streamed test is not enough).\n\n"
"If you know the endpoint streams tool calls correctly, set "
"STRIX_SKIP_TOOL_CALL_PROBE=1 to skip this check."
)
_TOOL_CONFIG_ERROR_MARKERS = (
"jinja",
"tool call parser",
"tool-call-parser",
"tool_choice",
"does not support tools",
"tools param",
"tool use is not supported",
)
class ToolCallingUnsupportedError(RuntimeError):
"""The endpoint cannot return structured tool calls over a streamed request."""
async def _noop_invoke(_ctx: Any, _args: str) -> str:
return "ok"
_PROBE_TOOL = FunctionTool(
name="strix_ready_check",
description="Report readiness. Call this to acknowledge you can use tools.",
params_json_schema={
"type": "object",
"additionalProperties": False,
"properties": {"status": {"type": "string"}},
"required": ["status"],
},
on_invoke_tool=_noop_invoke,
strict_json_schema=True,
)
_PROBE_SYSTEM = (
"You are a setup probe. You can only respond by calling the provided tool. "
"Do not produce any other output."
)
_PROBE_INPUT = 'Call the `strix_ready_check` tool now with {"status": "ok"}.'
def requires_tool_call_probe(model_name: str, settings: Settings) -> bool:
"""Only self-hosted / OpenAI-compatible routes, where the streaming leak happens."""
return model_name.startswith("ollama/") or bool(settings.llm.api_base)
def _is_tool_config_error(exc: Exception) -> bool:
text = str(exc).lower()
return any(marker in text for marker in _TOOL_CONFIG_ERROR_MARKERS)
async def _stream_saw_tool_call(
model_name: str, model_settings: ModelSettings, *, timeout: float | None
) -> bool:
model = StrixProvider().get_model(model_name)
async def _run() -> bool:
stream = model.stream_response(
system_instructions=_PROBE_SYSTEM,
input=_PROBE_INPUT,
model_settings=model_settings,
tools=[_PROBE_TOOL],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
previous_response_id=None,
conversation_id=None,
prompt=None,
)
# Drain the whole stream rather than returning early: the SDK wraps it in
# a span that must be closed in the context it was opened in.
saw_tool_call = False
async for event in stream:
item = getattr(event, "item", None)
if isinstance(item, ResponseFunctionToolCall):
saw_tool_call = True
response = getattr(event, "response", None)
if response is not None and any(
isinstance(out, ResponseFunctionToolCall)
for out in getattr(response, "output", []) or []
):
saw_tool_call = True
return saw_tool_call
if timeout is not None:
return await asyncio.wait_for(_run(), timeout=timeout)
return await _run()
async def probe_tool_calling(
model_name: str,
settings: Settings,
*,
request_timeout: float | None = None,
) -> None:
"""Fail fast if a streamed request to ``model_name`` yields no structured tool call.
No-op for hosted providers and when ``STRIX_SKIP_TOOL_CALL_PROBE`` is set.
"""
if settings.llm.skip_tool_call_probe or not requires_tool_call_probe(model_name, settings):
return
model_settings = ModelSettings(
parallel_tool_calls=False,
include_usage=True,
)
last_exc: Exception | None = None
for attempt in range(_PROBE_RETRIES + 1):
try:
saw_tool_call = await _stream_saw_tool_call(
model_name, model_settings, timeout=request_timeout
)
except ToolCallingUnsupportedError:
raise
except Exception as exc:
if _is_tool_config_error(exc):
logger.debug("Tool-call probe hit a tool-config error", exc_info=True)
raise ToolCallingUnsupportedError(_GUIDANCE) from exc
last_exc = exc
logger.debug(
"Tool-call probe attempt %d/%d failed transiently",
attempt + 1,
_PROBE_RETRIES + 1,
exc_info=True,
)
continue
if saw_tool_call:
logger.info("Tool-call probe passed for model %s", model_name)
return
raise ToolCallingUnsupportedError(_GUIDANCE)
# All attempts raised transient errors; surface the last one unchanged so the
# caller's existing connection-error handling reports it.
if last_exc is not None:
raise last_exc
+22
View File
@@ -33,6 +33,7 @@ from strix.config.models import (
)
from strix.core.inputs import DEFAULT_MAX_TURNS
from strix.core.paths import run_dir_for, runtime_state_dir
from strix.core.warmup import ToolCallingUnsupportedError, probe_tool_calling
from strix.interface.cli import run_cli
from strix.interface.tui import run_tui
from strix.interface.update_check import (
@@ -422,6 +423,27 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
)
logger.info("LLM warm-up succeeded for dedupe model %s", dedupe_model)
raw_model = (llm.model or "").strip()
await probe_tool_calling(raw_model, settings, request_timeout=llm.timeout)
except ToolCallingUnsupportedError as e:
logger.debug("Tool-call probe failed", exc_info=True)
error_text = Text()
error_text.append("TOOL CALLING NOT SUPPORTED", style="bold red")
error_text.append("\n\n", style="white")
error_text.append(str(e), style="white")
console.print("\n")
console.print(
Panel(
error_text,
title="[bold white]STRIX",
title_align="left",
border_style="red",
padding=(1, 2),
),
)
console.print()
sys.exit(1)
except Exception as e:
logger.debug("LLM warm-up failed", exc_info=True)
error_text = Text()
+74
View File
@@ -1,6 +1,7 @@
import json
import logging
import subprocess
import threading
from collections.abc import Callable
from datetime import UTC, datetime
from importlib.metadata import PackageNotFoundError, version
@@ -95,6 +96,8 @@ def get_global_report_state() -> Optional["ReportState"]:
def set_global_report_state(report_state: "ReportState") -> None:
global _global_report_state # noqa: PLW0603
_global_report_state = report_state
# New run: drop any streamed-cost entries a prior run left unconsumed.
streamed_openrouter_costs.clear()
class ReportState:
@@ -507,6 +510,72 @@ class ReportState:
self._sync_llm_usage_record()
def openrouter_stream_cost(usage: Any) -> float | None:
"""Total OpenRouter-reported cost from a raw stream ``usage`` block, or None.
Non-BYOK responses bill everything to ``usage.cost``. BYOK responses put the
OpenRouter fee in ``usage.cost`` (often 0) and the provider charge in
``usage.cost_details.upstream_inference_cost``, so BYOK totals sum the two.
"""
if not isinstance(usage, dict):
return None
total = 0.0
cost = usage.get("cost")
if isinstance(cost, int | float) and cost > 0:
total += float(cost)
if bool(usage.get("is_byok")):
details = usage.get("cost_details")
upstream = details.get("upstream_inference_cost") if isinstance(details, dict) else None
if isinstance(upstream, int | float) and upstream > 0:
total += float(upstream)
return total if total > 0 else None
def _response_id(completion_response: Any) -> str | None:
response_id = getattr(completion_response, "id", None)
if response_id is None and isinstance(completion_response, dict):
response_id = cast("dict[str, Any]", completion_response).get("id")
return response_id if isinstance(response_id, str) and response_id else None
class StreamedOpenRouterCosts:
"""Correlates OpenRouter's per-stream cost from the parser to the cost callback.
LiteLLM rebuilds streamed responses from token-only chunks and drops the
``usage.cost`` OpenRouter reports in its final stream chunk (its non-streamed
path preserves it; streaming snapshots hidden params at stream start). Every
scan streams, so the OpenRouter streaming handler (see strix.config.models)
records the cost here keyed by response id, and the callback takes it back out
for the matching rebuilt response. Entries are removed on read; ``clear()``
runs per scan so nothing accumulates across runs.
"""
def __init__(self) -> None:
self._costs: dict[str, float] = {}
self._lock = threading.Lock()
def remember(self, response_id: Any, usage: Any) -> None:
cost = openrouter_stream_cost(usage)
if cost is None or not (isinstance(response_id, str) and response_id):
return
with self._lock:
self._costs[response_id] = cost
def take(self, completion_response: Any) -> float | None:
response_id = _response_id(completion_response)
if response_id is None:
return None
with self._lock:
return self._costs.pop(response_id, None)
def clear(self) -> None:
with self._lock:
self._costs.clear()
streamed_openrouter_costs = StreamedOpenRouterCosts()
def litellm_cost_callback(
kwargs: Any,
completion_response: Any,
@@ -541,6 +610,11 @@ def litellm_cost_callback(
if cost is None:
cost = _usage_reported_cost(completion_response)
# Recover the exact OpenRouter cost the streaming handler stashed for this
# response — LiteLLM drops it from streamed usage, so nothing above sees it.
if cost is None:
cost = streamed_openrouter_costs.take(completion_response)
if cost is None:
cost = _estimate_response_cost(kwargs, completion_response)
+105 -2
View File
@@ -7,9 +7,25 @@ from unittest.mock import MagicMock, patch
import litellm
import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
from strix.config.models import _configure_litellm_compatibility
from strix.report.state import litellm_cost_callback
from strix.config.models import (
_configure_litellm_compatibility,
_install_openrouter_stream_cost_capture,
)
from strix.report.state import (
ReportState,
litellm_cost_callback,
openrouter_stream_cost,
set_global_report_state,
streamed_openrouter_costs,
)
@pytest.fixture(autouse=True)
def _clear_streamed_costs() -> None:
streamed_openrouter_costs.clear()
def test_streaming_logging_stays_enabled_for_cost_callback() -> None:
@@ -151,3 +167,90 @@ def test_cost_callback_records_nothing_when_no_cost_available() -> None:
litellm_cost_callback({"response_cost": None, "model": "x/y"}, response)
report_state.record_observed_llm_cost.assert_not_called()
def test_openrouter_stream_cost_extracts_plain_and_byok_totals() -> None:
assert openrouter_stream_cost({"cost": 0.003168}) == pytest.approx(0.003168)
assert openrouter_stream_cost(
{"cost": 0.01, "is_byok": True, "cost_details": {"upstream_inference_cost": 0.2}}
) == pytest.approx(0.21)
# Upstream cost is only added for BYOK responses.
assert openrouter_stream_cost(
{"cost": 0.05, "is_byok": False, "cost_details": {"upstream_inference_cost": 0.04}}
) == pytest.approx(0.05)
assert openrouter_stream_cost({"prompt_tokens": 10}) is None
assert openrouter_stream_cost(None) is None
def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> None:
report_state = MagicMock()
streamed_openrouter_costs.remember("gen-abc", {"cost": 0.42})
# LiteLLM strips cost from the rebuilt streamed usage; only the id survives.
response = SimpleNamespace(id="gen-abc", usage=SimpleNamespace(cost=None), _hidden_params={})
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=ValueError("unknown model")),
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.42)
# The entry is consumed so a later response cannot double-count it.
assert streamed_openrouter_costs.take(response) is None
def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None:
report_state = MagicMock()
streamed_openrouter_costs.remember("gen-xyz", {"cost": 0.9})
response = SimpleNamespace(
id="gen-xyz",
usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15),
_hidden_params={},
)
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", return_value=0.1) as estimate,
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.9)
estimate.assert_not_called()
def test_streamed_openrouter_costs_ignores_entries_without_cost() -> None:
streamed_openrouter_costs.remember("gen-none", {"prompt_tokens": 10})
streamed_openrouter_costs.remember("", {"cost": 0.5})
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-none")) is None
def test_streamed_openrouter_costs_cleared_on_new_run() -> None:
streamed_openrouter_costs.remember("gen-stale", {"cost": 0.7})
set_global_report_state(ReportState.__new__(ReportState))
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stale")) is None
def test_openrouter_stream_handler_records_cost() -> None:
_install_openrouter_stream_cost_capture()
# Resolve the config the way LiteLLM does in production so we prove the
# override is actually reachable through provider resolution, not just as a
# directly-constructed class.
config = ProviderConfigManager.get_provider_chat_config(
model="moonshotai/kimi-k3", provider=LlmProviders.OPENROUTER
)
assert config is not None
assert type(config).__name__ == "_StrixOpenrouterConfig"
handler = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True)
chunk = {
"id": "gen-stream",
"created": 1,
"model": "moonshotai/kimi-k3",
"choices": [{"index": 0, "delta": {"content": None}}],
"usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055},
}
handler.chunk_parser(chunk)
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx(
0.0035055
)
+158
View File
@@ -0,0 +1,158 @@
from __future__ import annotations
import types
from typing import TYPE_CHECKING, Any
import pytest
from openai.types.responses import ResponseFunctionToolCall
from strix.core import warmup
from strix.core.warmup import (
ToolCallingUnsupportedError,
probe_tool_calling,
requires_tool_call_probe,
)
if TYPE_CHECKING:
from collections.abc import AsyncIterator
def _settings(*, api_base: str | None = None, skip: bool = False) -> Any:
return types.SimpleNamespace(
llm=types.SimpleNamespace(api_base=api_base, skip_tool_call_probe=skip),
)
def _tool_call_event() -> Any:
return types.SimpleNamespace(
item=ResponseFunctionToolCall(
arguments='{"status": "ok"}',
call_id="call_1",
name="strix_ready_check",
type="function_call",
),
)
def _text_event() -> Any:
# A completed response whose output is a plain message, no tool call.
return types.SimpleNamespace(
item=None,
response=types.SimpleNamespace(output=[types.SimpleNamespace(type="message")]),
)
class _FakeModel:
def __init__(self, events: list[Any] | None = None, raises: Exception | None = None) -> None:
self._events = events or []
self._raises = raises
def stream_response(self, **_kwargs: Any) -> AsyncIterator[Any]:
events = self._events
raises = self._raises
async def _gen() -> AsyncIterator[Any]:
if raises is not None:
raise raises
for event in events:
yield event
return _gen()
def _patch_model(monkeypatch: pytest.MonkeyPatch, model: _FakeModel) -> None:
monkeypatch.setattr(
warmup, "StrixProvider", lambda: types.SimpleNamespace(get_model=lambda _m: model)
)
def test_requires_probe_only_for_custom_endpoints_and_ollama() -> None:
assert requires_tool_call_probe("openai/glm-5.2", _settings(api_base="http://x")) is True
assert requires_tool_call_probe("ollama/llama3", _settings()) is True
assert requires_tool_call_probe("openai/gpt-4o", _settings()) is False
assert requires_tool_call_probe("anthropic/claude", _settings()) is False
@pytest.mark.asyncio
async def test_probe_skipped_for_hosted_provider(monkeypatch: pytest.MonkeyPatch) -> None:
# Would raise if it tried to stream; gating must short-circuit first.
_patch_model(monkeypatch, _FakeModel(raises=RuntimeError("should not be called")))
await probe_tool_calling("openai/gpt-4o", _settings())
@pytest.mark.asyncio
async def test_probe_skipped_when_setting_disabled(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_model(monkeypatch, _FakeModel(raises=RuntimeError("should not be called")))
await probe_tool_calling("openai/glm-5.2", _settings(api_base="http://x", skip=True))
@pytest.mark.asyncio
async def test_probe_passes_on_streamed_tool_call(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_model(monkeypatch, _FakeModel(events=[_tool_call_event()]))
await probe_tool_calling("openai/glm-5.2", _settings(api_base="http://x"))
@pytest.mark.asyncio
async def test_probe_passes_when_tool_call_only_in_final_response(
monkeypatch: pytest.MonkeyPatch,
) -> None:
completed = types.SimpleNamespace(
item=None,
response=types.SimpleNamespace(
output=[
ResponseFunctionToolCall(
arguments="{}",
call_id="c",
name="strix_ready_check",
type="function_call",
)
]
),
)
_patch_model(monkeypatch, _FakeModel(events=[completed]))
await probe_tool_calling("openai/glm-5.2", _settings(api_base="http://x"))
@pytest.mark.asyncio
async def test_probe_aborts_when_only_text_streamed(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_model(monkeypatch, _FakeModel(events=[_text_event()]))
with pytest.raises(ToolCallingUnsupportedError):
await probe_tool_calling("openai/glm-5.2", _settings(api_base="http://x"))
@pytest.mark.asyncio
async def test_probe_aborts_on_tool_config_error(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_model(monkeypatch, _FakeModel(raises=RuntimeError("tools param requires --jinja flag")))
with pytest.raises(ToolCallingUnsupportedError):
await probe_tool_calling("ollama/llama3", _settings(api_base="http://x"))
@pytest.mark.asyncio
async def test_probe_retries_transient_then_passes(monkeypatch: pytest.MonkeyPatch) -> None:
calls = {"n": 0}
good = _FakeModel(events=[_tool_call_event()])
class _Flaky:
def stream_response(self, **kwargs: Any) -> AsyncIterator[Any]:
calls["n"] += 1
if calls["n"] == 1:
async def _boom() -> AsyncIterator[Any]:
for _ in range(0): # make this a generator without an unreachable yield
yield None
raise ConnectionError("transient")
return _boom()
return good.stream_response(**kwargs)
_patch_model(monkeypatch, _Flaky()) # type: ignore[arg-type]
await probe_tool_calling("openai/glm-5.2", _settings(api_base="http://x"))
assert calls["n"] == 2
@pytest.mark.asyncio
async def test_probe_surfaces_persistent_transient_error(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_model(monkeypatch, _FakeModel(raises=ConnectionError("down")))
with pytest.raises(ConnectionError):
await probe_tool_calling("openai/glm-5.2", _settings(api_base="http://x"))