diff --git a/strix/config/models.py b/strix/config/models.py index e33f30604..6f0162a00 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -792,6 +792,16 @@ def _install_openrouter_stream_cost_capture() -> None: json_mode=json_mode, ) + def transform_response(self, *args: Any, **kwargs: Any) -> Any: + # Non-streamed replies (LLM_DISABLE_STREAMING) skip the chunk parser. + response = super().transform_response(*args, **kwargs) + raw_response = kwargs.get("raw_response", args[1] if len(args) > 1 else None) + with contextlib.suppress(Exception): + body = raw_response.json() # type: ignore[union-attr] + if body.get("usage"): + record_openrouter_provider(body.get("provider"), body["usage"]) + return response + def transform_request(self, *args: Any, **kwargs: Any) -> dict[str, Any]: # Pin each agent's calls to one upstream provider so its prompt cache # survives between turns. diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index 4c0ac9223..9d7ea6db4 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -5,8 +5,9 @@ from __future__ import annotations import uuid from types import SimpleNamespace from typing import TYPE_CHECKING, Any -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch +import httpx import litellm import pytest from litellm.types.utils import LlmProviders @@ -267,7 +268,7 @@ def test_openrouter_stream_handler_records_cost() -> None: ) -def test_openrouter_stream_handler_tallies_provider() -> None: +def test_openrouter_tallies_provider() -> None: _install_openrouter_stream_cost_capture() config = ProviderConfigManager.get_provider_chat_config( model="z-ai/glm-5.3", provider=LlmProviders.OPENROUTER @@ -292,10 +293,26 @@ def test_openrouter_stream_handler_tallies_provider() -> None: "usage": usage, } ) + # Non-streamed replies (LLM_DISABLE_STREAMING) carry the same fields. + reply = { + "choices": [{"message": {"role": "assistant"}}], + "provider": "Together", + "usage": usage, + } + config.transform_response( + "z-ai/glm-5.3", + httpx.Response(200, json=reply), + litellm.ModelResponse(), + MagicMock(), + {}, + [], + {}, + {}, + None, + ) - report_state.record_llm_provider.assert_called_once_with( - "Together", agent_id=None, input_tokens=1000, cached_tokens=900, cost=0.002 - ) + tally = call("Together", agent_id=None, input_tokens=1000, cached_tokens=900, cost=0.002) + assert report_state.record_llm_provider.call_args_list == [tally, tally] def test_provider_tally_survives_run_record_round_trip() -> None: