diff --git a/src/ai_client.py b/src/ai_client.py index cdd88a95..24922a32 100644 --- a/src/ai_client.py +++ b/src/ai_client.py @@ -2025,6 +2025,7 @@ def _send_gemini_cli(md_content: str, user_message: str, base_dir: str, stream_callback: Optional[Callable[[str], None]] = None, patch_callback: Optional[Callable[[str, str], Optional[str]]] = None) -> Result[str]: from src.openai_compatible import OpenAICompatibleRequest, NormalizedResponse + from src.openai_schemas import UsageStats """ [C: src/ai_server.py:_handle_send] Functional Purpose: Sends requests to Gemini via the headless Gemini CLI subprocess adapter. @@ -2051,7 +2052,7 @@ def _send_gemini_cli(md_content: str, user_message: str, base_dir: str, def _send(r_idx: int) -> NormalizedResponse: if adapter is None: - return NormalizedResponse(text="(adapter unavailable)", tool_calls=[], usage_input_tokens=0, usage_output_tokens=0, usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=None) + return NormalizedResponse(text="(adapter unavailable)", tool_calls=[], usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None) send_result = _send_cli_round_result(r_idx, adapter, payload, safety_settings, sys_instr, stream_callback) if not send_result.ok: raise cast(Exception, send_result.errors[0].original) from None @@ -2085,7 +2086,7 @@ def _send_gemini_cli(md_content: str, user_message: str, base_dir: str, "kind": "history_add", "payload": {"role": "AI", "content": txt} }) - return NormalizedResponse(text=txt, tool_calls=calls, usage_input_tokens=usage.get("prompt_tokens", 0), usage_output_tokens=usage.get("completion_tokens", 0), usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=resp_data) + return NormalizedResponse(text=txt, tool_calls=calls, usage=UsageStats(input_tokens=usage.get("prompt_tokens", 0), output_tokens=usage.get("completion_tokens", 0), cache_read_tokens=0, cache_creation_tokens=0), raw_response=resp_data) def _pre_dispatch(r_idx: int, calls: list[Metadata]) -> list[Metadata]: nonlocal payload, cumulative_tool_bytes, file_items @@ -2569,7 +2570,7 @@ def _send_grok(md_content: str, user_message: str, base_dir: str, Runs synchronously in the caller thread; synchronizes Grok history using _grok_history_lock. """ from src.openai_compatible import OpenAICompatibleRequest, _classify_openai_compatible_error - from src.openai_schemas import ChatMessage + from src.openai_schemas import ChatMessage, UsageStats try: client = _ensure_grok_client() tools: list[Metadata] | None = _get_deepseek_tools() or None diff --git a/src/openai_schemas.py b/src/openai_schemas.py index 0058a4be..76dd5e2e 100644 --- a/src/openai_schemas.py +++ b/src/openai_schemas.py @@ -16,7 +16,7 @@ CONVENTION: 1-space indentation. NO COMMENTS. """ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Any, Callable, Optional from src.type_aliases import JsonValue @@ -72,35 +72,12 @@ class UsageStats: cache_creation_tokens: int = 0 -@dataclass(frozen=True, init=False) +@dataclass(frozen=True) class NormalizedResponse: text: str - tool_calls: tuple[ToolCall, ...] - usage: UsageStats - raw_response: Any - - def __init__( - self, - text: str, - tool_calls: tuple[ToolCall, ...] = (), - usage: UsageStats | None = None, - raw_response: Any = None, - usage_input_tokens: int | None = None, - usage_output_tokens: int | None = None, - usage_cache_read_tokens: int | None = None, - usage_cache_creation_tokens: int | None = None, - ) -> None: - if usage is None: - usage = UsageStats( - input_tokens=usage_input_tokens if usage_input_tokens is not None else 0, - output_tokens=usage_output_tokens if usage_output_tokens is not None else 0, - cache_read_tokens=usage_cache_read_tokens if usage_cache_read_tokens is not None else 0, - cache_creation_tokens=usage_cache_creation_tokens if usage_cache_creation_tokens is not None else 0, - ) - object.__setattr__(self, "text", text) - object.__setattr__(self, "tool_calls", tool_calls) - object.__setattr__(self, "usage", usage) - object.__setattr__(self, "raw_response", raw_response) + tool_calls: tuple[ToolCall, ...] = () + usage: UsageStats = field(default_factory=lambda: UsageStats(input_tokens=0, output_tokens=0)) + raw_response: Any = None def to_legacy_dict(self) -> JsonValue: return { diff --git a/tests/test_ai_client_tool_loop.py b/tests/test_ai_client_tool_loop.py index eb576dc6..e6f1b4c8 100644 --- a/tests/test_ai_client_tool_loop.py +++ b/tests/test_ai_client_tool_loop.py @@ -18,6 +18,7 @@ from unittest.mock import MagicMock, patch import pytest from src.result_types import Result from src.openai_compatible import NormalizedResponse, OpenAICompatibleRequest +from src.openai_schemas import UsageStats from src.ai_client import run_with_tool_loop from src.vendor_capabilities import VendorCapabilities @@ -28,8 +29,7 @@ def caps() -> VendorCapabilities: def _make_normalized_response(text: str = "ok", tool_calls: list[dict[str, Any]] | None = None) -> Result[NormalizedResponse]: return Result(data=NormalizedResponse( text=text, tool_calls=tool_calls or [], - usage_input_tokens=10, usage_output_tokens=5, - usage_cache_read_tokens=0, usage_cache_creation_tokens=0, + usage=UsageStats(input_tokens=10, output_tokens=5, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None, )) diff --git a/tests/test_ai_client_tool_loop_builder.py b/tests/test_ai_client_tool_loop_builder.py index e7fae125..76a3437e 100644 --- a/tests/test_ai_client_tool_loop_builder.py +++ b/tests/test_ai_client_tool_loop_builder.py @@ -8,6 +8,7 @@ from __future__ import annotations from typing import Any from unittest.mock import MagicMock, patch from src.openai_compatible import NormalizedResponse, OpenAICompatibleRequest +from src.openai_schemas import UsageStats from src.ai_client import run_with_tool_loop from src.result_types import Result from src.vendor_capabilities import VendorCapabilities @@ -15,8 +16,7 @@ from src.vendor_capabilities import VendorCapabilities def _make_normalized_response(text: str = "ok", tool_calls: list[dict[str, Any]] | None = None) -> NormalizedResponse: return NormalizedResponse( text=text, tool_calls=tool_calls or [], - usage_input_tokens=10, usage_output_tokens=5, - usage_cache_read_tokens=0, usage_cache_creation_tokens=0, + usage=UsageStats(input_tokens=10, output_tokens=5, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None, ) diff --git a/tests/test_ai_client_tool_loop_send_func.py b/tests/test_ai_client_tool_loop_send_func.py index d46501f9..c4df65bd 100644 --- a/tests/test_ai_client_tool_loop_send_func.py +++ b/tests/test_ai_client_tool_loop_send_func.py @@ -7,14 +7,14 @@ from __future__ import annotations from typing import Any from unittest.mock import MagicMock, patch from src.openai_compatible import NormalizedResponse +from src.openai_schemas import UsageStats from src.ai_client import run_with_tool_loop from src.vendor_capabilities import VendorCapabilities def _make_normalized_response(text: str = "ok", tool_calls: list[dict[str, Any]] | None = None) -> NormalizedResponse: return NormalizedResponse( text=text, tool_calls=tool_calls or [], - usage_input_tokens=10, usage_output_tokens=5, - usage_cache_read_tokens=0, usage_cache_creation_tokens=0, + usage=UsageStats(input_tokens=10, output_tokens=5, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None, ) diff --git a/tests/test_ai_loop_regressions_20260614.py b/tests/test_ai_loop_regressions_20260614.py index 08fe550e..0c0a6e61 100644 --- a/tests/test_ai_loop_regressions_20260614.py +++ b/tests/test_ai_loop_regressions_20260614.py @@ -19,6 +19,7 @@ from src import ai_client from src import thinking_parser from src.gui_2 import App from src.events import UserRequestEvent +from src.openai_schemas import UsageStats from src.result_types import Result, ErrorInfo, ErrorKind @@ -206,10 +207,7 @@ def test_fr3_minimax_thinking_in_returned_text() -> None: return Result(data=MagicMock( text="The final answer is 42", tool_calls=[], - usage_input_tokens=0, - usage_output_tokens=0, - usage_cache_read_tokens=0, - usage_cache_creation_tokens=0, + usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=fake_raw, )) diff --git a/tests/test_grok_provider.py b/tests/test_grok_provider.py index b1dc6b3b..0d54c203 100644 --- a/tests/test_grok_provider.py +++ b/tests/test_grok_provider.py @@ -30,10 +30,11 @@ def test_grok_2_vision_supports_image() -> None: def test_grok_web_search_adds_search_parameters_to_extra_body() -> None: """caps.web_search=True should populate search_parameters.mode=auto in extra_body.""" from src import openai_compatible as oc + from src.openai_schemas import UsageStats captured_kwargs: list[dict] = [] def _fake_send(client, request, *, capabilities): captured_kwargs.append({"extra_body": request.extra_body, "model": request.model}) - return MagicMock(text="ok", tool_calls=[], usage_input_tokens=0, usage_output_tokens=0, usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=None) + return MagicMock(text="ok", tool_calls=[], usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None) with patch.object(oc, "send_openai_compatible", side_effect=_fake_send), \ patch("src.ai_client._ensure_grok_client", return_value=MagicMock()), \ patch("src.ai_client._get_deepseek_tools", return_value=[]): @@ -43,12 +44,13 @@ def test_grok_web_search_adds_search_parameters_to_extra_body() -> None: def test_grok_x_search_adds_x_source_to_extra_body() -> None: """caps.x_search=True should add sources=[{type:x}] to search_parameters.""" from src import openai_compatible as oc + from src.openai_schemas import UsageStats captured_kwargs: list[dict] = [] def _fake_send(client, request, *, capabilities): captured_kwargs.append({"extra_body": request.extra_body}) - return MagicMock(text="ok", tool_calls=[], usage_input_tokens=0, usage_output_tokens=0, usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=None) + return MagicMock(text="ok", tool_calls=[], usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None) with patch.object(oc, "send_openai_compatible", side_effect=_fake_send), \ patch("src.ai_client._ensure_grok_client", return_value=MagicMock()), \ patch("src.ai_client._get_deepseek_tools", return_value=[]): ai_client._send_grok("system", "user", ".", None, "", False, None, None, None) - assert captured_kwargs[0]["extra_body"]["search_parameters"]["sources"] == [{"type": "x"}] \ No newline at end of file + assert captured_kwargs[0]["extra_body"]["search_parameters"]["sources"] == [{"type": "x"}] \ No newline at end of file diff --git a/tests/test_minimax_provider.py b/tests/test_minimax_provider.py index f37f5afe..685b2a4f 100644 --- a/tests/test_minimax_provider.py +++ b/tests/test_minimax_provider.py @@ -37,10 +37,11 @@ def test_minimax_credentials_template() -> None: def test_minimax_reasoning_extractor_used_when_caps_reasoning_true() -> None: """caps.reasoning=True (M2.5/M2.7) should pass the reasoning_extractor to run_with_tool_loop.""" from src import openai_compatible as oc + from src.openai_schemas import UsageStats captured_kwargs: list[dict] = [] def _fake_send(client, request, *, capabilities): captured_kwargs.append({"model": request.model}) - return MagicMock(text="ok", tool_calls=[], usage_input_tokens=0, usage_output_tokens=0, usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=None) + return MagicMock(text="ok", tool_calls=[], usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None) from src.vendor_capabilities import register, VendorCapabilities register(VendorCapabilities(vendor='minimax', model='MiniMax-M2.5', reasoning=True)) with patch.object(oc, "send_openai_compatible", side_effect=_fake_send), \ @@ -52,17 +53,18 @@ def test_minimax_reasoning_extractor_used_when_caps_reasoning_true() -> None: def test_minimax_reasoning_extractor_omitted_when_caps_reasoning_false() -> None: """caps.reasoning=False (M2/M2.1) should NOT pass the reasoning_extractor (avoid useless getattr).""" from src import openai_compatible as oc + from src.openai_schemas import UsageStats from src.vendor_capabilities import register, VendorCapabilities register(VendorCapabilities(vendor='minimax', model='MiniMax-M2', reasoning=False)) captured_kwargs: list[dict] = [] def _fake_send(client, request, *, capabilities): captured_kwargs.append({"model": request.model}) - return MagicMock(text="ok", tool_calls=[], usage_input_tokens=0, usage_output_tokens=0, usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=None) + return MagicMock(text="ok", tool_calls=[], usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None) with patch.object(oc, "send_openai_compatible", side_effect=_fake_send), \ patch("src.ai_client._ensure_minimax_client", return_value=MagicMock()), \ patch("src.ai_client._get_deepseek_tools", return_value=[]): ai_client._send_minimax("system", "user", ".", None, "", False, None, None, None) - assert len(captured_kwargs) >= 1 + assert len(captured_kwargs) >= 1 def test_minimax_ensure_client_instantiation() -> None: """Verify that _ensure_minimax_client instantiates the OpenAI client with correct credentials and base URL.""" diff --git a/tests/test_openai_compatible.py b/tests/test_openai_compatible.py index 5c5344a9..327aa718 100644 --- a/tests/test_openai_compatible.py +++ b/tests/test_openai_compatible.py @@ -86,6 +86,7 @@ def test_error_classification_429_to_rate_limit(caps: VendorCapabilities) -> Non def test_normalized_response_is_frozen_dataclass() -> None: from dataclasses import FrozenInstanceError - r = NormalizedResponse(text="x", tool_calls=[], usage_input_tokens=0, usage_output_tokens=0, usage_cache_read_tokens=0, usage_cache_creation_tokens=0, raw_response=None) + from src.openai_schemas import UsageStats + r = NormalizedResponse(text="x", tool_calls=[], usage=UsageStats(input_tokens=0, output_tokens=0, cache_read_tokens=0, cache_creation_tokens=0), raw_response=None) with pytest.raises(FrozenInstanceError): r.text = "y"