From 350b88d9ff91e2476f0bfaaac3f4c53be22170f4 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:40:52 -0700 Subject: [PATCH] fix(responses): log and bill every round of non-stream MCP auto-execute Both model rounds run as internal sub-calls on the request's logging object, and the outer call logs the response the client received with usage, cost, and cost breakdown summed across the rounds. The internal-call gate now also covers the sync success callbacks, so an inner round no longer takes their first-wins slot --- litellm/_internal_context.py | 10 + litellm/rag/main.py | 30 +-- litellm/responses/main.py | 49 ++-- litellm/responses/mcp/round_billing.py | 83 ++++++ litellm/utils.py | 17 +- .../mcp/test_aresponses_api_with_mcp.py | 242 +++++++++++++++++- .../unit/responses/mcp/test_round_billing.py | 113 ++++++++ 7 files changed, 490 insertions(+), 54 deletions(-) create mode 100644 litellm/responses/mcp/round_billing.py create mode 100644 tests/unit/responses/mcp/test_round_billing.py diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py index 389add8ed0f..2680accb49e 100644 --- a/litellm/_internal_context.py +++ b/litellm/_internal_context.py @@ -24,6 +24,16 @@ _billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", d _post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False) +@contextmanager +def internal_sub_call() -> Generator[None]: + """A nested call the parent call logs and bills, so the parent must fold this call's cost into its own.""" + token: Final = is_internal_call.set(True) + try: + yield + finally: + is_internal_call.reset(token) + + @contextmanager def post_response_phase() -> Generator[None]: """Work the caller no longer waits for (success callbacks, response-cache writes), including tasks it spawns.""" diff --git a/litellm/rag/main.py b/litellm/rag/main.py index 1a5301f0579..59c3c20c396 100644 --- a/litellm/rag/main.py +++ b/litellm/rag/main.py @@ -11,8 +11,7 @@ __all__ = ["aingest", "aquery", "ingest", "query"] import asyncio import contextvars -from collections.abc import Coroutine, Iterator, Mapping -from contextlib import contextmanager +from collections.abc import Coroutine, Mapping from functools import partial from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final @@ -20,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Final import httpx import litellm -from litellm._internal_context import is_internal_call +from litellm._internal_context import internal_sub_call from litellm.cost_calculator import vector_store_search_cost from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.rag.ingestion.base_ingestion import BaseRAGIngestion @@ -204,25 +203,6 @@ async def aingest( ) -@contextmanager -def _suppressed_sub_call_billing() -> Iterator[None]: - """ - Suppress a sub-call's own billing event so the parent aquery event bills it. - - Every suppressed sub-call's cost must be folded into the parent event: - into the response's hidden response_cost on the non-streaming path, or via - the logging object's additional_response_cost on the streaming path (the - streamed cost is computed from assembled chunks after this pipeline - returns, so there is no response object to fold into here). - """ - previous: Final = is_internal_call.get() - is_internal_call.set(True) - try: - yield - finally: - is_internal_call.set(previous) - - async def _execute_query_pipeline( model: str, messages: list[AllMessageValues], @@ -264,7 +244,7 @@ async def _execute_query_pipeline( forwarded_search_params: Final = MappingProxyType( {**provider_search_params, **kwargs, **filter_search_params, **store_search_params} ) - with _suppressed_sub_call_billing(): + with internal_sub_call(): search_response: Final = await litellm.vector_stores.asearch( vector_store_id=retrieval_config["vector_store_id"], query=query_text, @@ -294,7 +274,7 @@ async def _execute_query_pipeline( if rerank and rerank.get("enabled"): documents: Final = RAGQuery.extract_documents_from_search(search_response) if documents: - with _suppressed_sub_call_billing(): + with internal_sub_call(): rerank_response = await litellm.arerank( model=rerank["model"], query=query_text, @@ -312,7 +292,7 @@ async def _execute_query_pipeline( modified_messages: Final = messages[:-1] + [context_message] + [messages[-1]] # Use router if available to properly resolve virtual model names - with _suppressed_sub_call_billing(): + with internal_sub_call(): if router is not None: response = await router.acompletion( model=model, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index f1ed8d3e9b3..c7a0cb65a6b 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -13,6 +13,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError from typing_extensions import assert_never import litellm +from litellm._internal_context import internal_sub_call from litellm._logging import verbose_logger from litellm.completion_extras.litellm_responses_transformation.transformation import ( LiteLLMResponsesTransformationHandler, @@ -34,6 +35,7 @@ from litellm.responses.litellm_completion_transformation.handler import ( LiteLLMCompletionTransformationHandler, ) from litellm.responses.mcp.request_context import MCPRequestContext +from litellm.responses.mcp.round_billing import billed_for_every_round, summed_cost_breakdown from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ( PromptObject, @@ -306,13 +308,16 @@ async def aresponses_api_with_mcp( ######################################################### # Make initial response API call ######################################################### - response: Final = await aresponses( - input=input, - model=model, - tools=all_tools, - previous_response_id=previous_response_id, - **initial_call_params, - ) + with internal_sub_call(): + response: Final = await aresponses( + input=input, + model=model, + tools=all_tools, + previous_response_id=previous_response_id, + **initial_call_params, + ) + logging_obj: Final[object] = kwargs.get("litellm_logging_obj") + first_round_breakdown: Final = logging_obj.cost_breakdown if isinstance(logging_obj, LiteLLMLoggingObj) else None verbose_logger.debug("Initial response %s", response) @@ -375,13 +380,14 @@ async def aresponses_api_with_mcp( tool_calls=tool_calls, tool_results=tool_results ) - final_response = await LiteLLM_Proxy_MCP_Handler._make_follow_up_call( - follow_up_input=follow_up_input, - model=model, - all_tools=all_tools, - response_id=previous_response_id if persistence_disabled else response.id, - **follow_up_call_params, - ) + with internal_sub_call(): + final_response = await LiteLLM_Proxy_MCP_Handler._make_follow_up_call( + follow_up_input=follow_up_input, + model=model, + all_tools=all_tools, + response_id=previous_response_id if persistence_disabled else response.id, + **follow_up_call_params, + ) # If streaming and we have tool execution events, wrap the response if ( @@ -414,10 +420,17 @@ async def aresponses_api_with_mcp( request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs), raw_headers=discovery_raw_headers, ) - final_response = LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response( - response=final_response, - mcp_tools_fetched=mcp_tools_for_output, - tool_results=tool_results, + if isinstance(logging_obj, LiteLLMLoggingObj): + logging_obj.cost_breakdown = summed_cost_breakdown( + first_round_breakdown, logging_obj.cost_breakdown + ) + return billed_for_every_round( + first=response, + final=LiteLLM_Proxy_MCP_Handler._add_mcp_output_elements_to_response( + response=final_response, + mcp_tools_fetched=mcp_tools_for_output, + tool_results=tool_results, + ), ) return final_response diff --git a/litellm/responses/mcp/round_billing.py b/litellm/responses/mcp/round_billing.py new file mode 100644 index 00000000000..efbe65d1ee1 --- /dev/null +++ b/litellm/responses/mcp/round_billing.py @@ -0,0 +1,83 @@ +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + +from pydantic import TypeAdapter +from typing_extensions import TypeIs + +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse +from litellm.types.utils import CostBreakdown + +_RATE_KEYS: Final = frozenset({"discount_percent", "margin_percent"}) +_COST_BREAKDOWN: Final = TypeAdapter(CostBreakdown) +_HIDDEN_PARAMS: Final = TypeAdapter(Mapping[str, object]) + + +def _is_str_keyed( + value: object, +) -> TypeIs[Mapping[str, object]]: # guard-ok: model dumps and CostBreakdown have str keys + return isinstance(value, Mapping) + + +def _summed(first: object, final: object) -> object: + if first is None: + return final + if final is None: + return first + if isinstance(first, bool) or isinstance(final, bool): + return final + if isinstance(first, (int, float)) and isinstance(final, (int, float)): + return first + final + if _is_str_keyed(first) and _is_str_keyed(final): + return _summed_mapping(first, final) + return final + + +def _summed_entry(key: str, first: Mapping[str, object], final: Mapping[str, object]) -> object: + if key in _RATE_KEYS: + return final[key] if key in final else first[key] + return _summed(first.get(key), final.get(key)) + + +def _summed_mapping(first: Mapping[str, object], final: Mapping[str, object]) -> Mapping[str, object]: + keys: Final = (*first, *(key for key in final if key not in first)) + return MappingProxyType({key: _summed_entry(key, first, final) for key in keys}) + + +def merged_round_usage(first: ResponseAPIUsage | None, final: ResponseAPIUsage | None) -> ResponseAPIUsage | None: + if first is None or final is None: + return final if first is None else first + return ResponseAPIUsage.model_validate(_summed_mapping(first.model_dump(), final.model_dump())) + + +def summed_cost_breakdown(first: CostBreakdown | None, final: CostBreakdown | None) -> CostBreakdown | None: + """None when either round was priced without a breakdown, since a partial one would disagree with the billed cost.""" + if first is None or final is None or first is final: + return None + return _COST_BREAKDOWN.validate_python(_summed_mapping(first, final)) + + +def _hidden_params(response: ResponsesAPIResponse) -> Mapping[str, object]: + return _HIDDEN_PARAMS.validate_python(getattr(response, "_hidden_params", None)) + + +def billed_for_every_round(first: ResponsesAPIResponse, final: ResponsesAPIResponse) -> ResponsesAPIResponse: + """The final round's response carrying the usage and cost of both rounds, so the one logged row bills the request. + + Costs are summed per round rather than repriced from the summed usage, because a round's price can depend on its + own size (tiered rates) or on charges the usage does not show (built-in tools, provider-reported costs). + """ + final_hidden_params: Final = _hidden_params(final) + first_cost: Final = _hidden_params(first).get("response_cost") + final_cost: Final = final_hidden_params.get("response_cost") + total_cost: Final = ( + first_cost + final_cost + if isinstance(first_cost, (int, float)) and isinstance(final_cost, (int, float)) + else None + ) + billed: Final = final.model_copy(update=MappingProxyType({"usage": merged_round_usage(first.usage, final.usage)})) + billed._hidden_params = { # pyright: ignore[reportPrivateUsage] # no public accessor # mutable-ok: the outer call writes into this dict + **final_hidden_params, + "response_cost": total_cost, + } + return billed diff --git a/litellm/utils.py b/litellm/utils.py index 13a46840431..1d50f70a277 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1323,15 +1323,16 @@ def _dispatch_success_logging( is_completion_with_fallbacks: bool, is_litellm_internal_call: bool, ) -> None: - if not is_litellm_internal_call: - _schedule_async_success_logging( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) + if is_litellm_internal_call: + return + _schedule_async_success_logging( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) logging_obj.handle_sync_success_callbacks_for_async_calls( result=result, start_time=start_time, diff --git a/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py b/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py index de0dc78af43..a91d1a5fb9d 100644 --- a/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py +++ b/tests/unit/responses/mcp/test_aresponses_api_with_mcp.py @@ -1,11 +1,21 @@ -import pytest -from mcp.types import Tool as MCPTool -from typing import List, Any, cast +import asyncio +import json +from collections.abc import Callable, Mapping, Sequence +from datetime import datetime +from types import MappingProxyType +from typing import Any, Final, List, cast from unittest.mock import AsyncMock, patch +import httpx +import pytest +import respx +from mcp.types import Tool as MCPTool + # Import required modules import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.responses.mcp.litellm_proxy_mcp_handler import LiteLLM_Proxy_MCP_Handler from litellm.types.llms.openai import ( ResponsesAPIResponse, @@ -13,6 +23,7 @@ from litellm.types.llms.openai import ( OpenAIMcpServerTool, ToolParam, ) +from litellm.types.utils import StandardLoggingPayload class MockUserAPIKeyAuth: @@ -1275,3 +1286,228 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e(): "tools_per_call": [len(tools) for tools in llm_call_tools], "duplicate_tools_found": False, } + + +_ROUND_ONE_BODY: Final = MappingProxyType( + { + "id": "resp_round_one", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.6", + "output": ( + MappingProxyType( + { + "type": "function_call", + "id": "fc_round_one", + "call_id": "call_echo", + "name": "echo", + "arguments": '{"value": "PING"}', + "status": "completed", + } + ), + ), + "usage": MappingProxyType( + { + "input_tokens": 71, + "input_tokens_details": MappingProxyType({"cached_tokens": 32}), + "output_tokens": 19, + "output_tokens_details": MappingProxyType({"reasoning_tokens": 0}), + "total_tokens": 90, + } + ), + } +) +_ROUND_TWO_BODY: Final = MappingProxyType( + { + "id": "resp_round_two", + "object": "response", + "created_at": 2, + "status": "completed", + "model": "gpt-5.6", + "output": ( + MappingProxyType( + { + "type": "message", + "id": "msg_round_two", + "status": "completed", + "role": "assistant", + "content": (MappingProxyType({"type": "output_text", "text": "PING", "annotations": ()}),), + } + ), + ), + "usage": MappingProxyType( + { + "input_tokens": 149, + "input_tokens_details": MappingProxyType({"cached_tokens": 64}), + "output_tokens": 5, + "output_tokens_details": MappingProxyType({"reasoning_tokens": 0}), + "total_tokens": 154, + } + ), + } +) + + +def _round_cost(body: Mapping[str, object]) -> float: + return litellm.completion_cost( + completion_response=ResponsesAPIResponse.model_validate(body), + model="gpt-5.6", + custom_llm_provider="openai", + call_type="aresponses", + ) + + +class _RoundRecorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.logged: tuple[StandardLoggingPayload, ...] = () + self.hooked: tuple[object, ...] = () + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + self.logged = (*self.logged, cast(StandardLoggingPayload, kwargs["standard_logging_object"])) + + def logging_hook(self, kwargs: dict, result: Any, call_type: str) -> tuple[dict, Any]: + self.hooked = (*self.hooked, result) + return kwargs, result + + +def _no_op_sync_callback( + kwargs: Mapping[str, object], completion_response: object, start_time: datetime, end_time: datetime +) -> None: + return None + + +async def _settle(condition: Callable[[], bool]) -> None: + for _ in range(200): + if condition(): + break + await asyncio.sleep(0.01) + await asyncio.sleep(0.2) + + +async def _auto_execute_echo( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, logging_obj: Logging +) -> ResponsesAPIResponse: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + litellm.in_memory_llm_clients_cache.flush_cache() + respx_mock.post(url__regex=r"https://mcp-rounds\.example\.com/.*responses").mock( + side_effect=[ + httpx.Response(200, content=json.dumps(_ROUND_ONE_BODY, default=dict)), + httpx.Response(200, content=json.dumps(_ROUND_TWO_BODY, default=dict)), + ] + ) + echo_tool: Final = MCPTool.model_validate( + {"name": "echo", "inputSchema": {"type": "object", "properties": {"value": {"type": "string"}}}}, + by_name=False, + ) + with ( + patch.object( + LiteLLM_Proxy_MCP_Handler, + "_process_mcp_tools_without_openai_transform", + AsyncMock(return_value=([echo_tool], {"echo": "demo"})), + ), + patch.object( + LiteLLM_Proxy_MCP_Handler, + "_execute_tool_calls", + AsyncMock(return_value=[{"tool_call_id": "call_echo", "result": "PING", "name": "echo"}]), + ), + ): + response: Final = await litellm.aresponses( + model="openai/gpt-5.6", + input="Call echo with PING", + tools=[{"type": "mcp", "server_url": "litellm_proxy/mcp/demo", "require_approval": "never"}], + api_key="sk-mcp-rounds", + api_base="https://mcp-rounds.example.com/v1", + litellm_logging_obj=logging_obj, + ) + assert isinstance(response, ResponsesAPIResponse) + return response + + +def _rounds_logging_obj(recorder: _RoundRecorder, call_id: str) -> Logging: + return Logging( + model="openai/gpt-5.6", + messages=[{"role": "user", "content": "Call echo with PING"}], + stream=False, + call_type="aresponses", + start_time=datetime.now(), + litellm_call_id=call_id, + function_id=call_id, + dynamic_async_success_callbacks=[recorder], + dynamic_success_callbacks=[recorder, _no_op_sync_callback], + ) + + +def _output_types(output: Sequence[object]) -> tuple[object, ...]: + return tuple(item.get("type") if isinstance(item, Mapping) else getattr(item, "type", None) for item in output) + + +def _assert_one_row_bills_both_rounds(logged: tuple[StandardLoggingPayload, ...], client: ResponsesAPIResponse) -> None: + assert len(logged) == 1 + row: Final = logged[0] + both_rounds: Final = _round_cost(_ROUND_ONE_BODY) + _round_cost(_ROUND_TWO_BODY) + assert row["id"] == client.id + assert _output_types(row["response"]["output"]) == _output_types(client.output) + assert "message" in _output_types(client.output) + assert "function_call" not in _output_types(client.output) + assert (row["prompt_tokens"], row["completion_tokens"]) == (71 + 149, 19 + 5) + assert row["metadata"]["usage_object"]["prompt_tokens_details"]["cached_tokens"] == 32 + 64 + assert row["response_cost"] == pytest.approx(both_rounds) + assert row["cost_breakdown"]["total_cost"] == pytest.approx(both_rounds) + assert client._hidden_params["response_cost"] == pytest.approx(both_rounds) + assert client.usage is not None + assert (client.usage.input_tokens, client.usage.output_tokens) == (71 + 149, 19 + 5) + + +@pytest.mark.asyncio +async def test_non_stream_auto_execute_logs_the_client_response_billed_for_both_rounds( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + recorder: Final = _RoundRecorder() + response: Final = await _auto_execute_echo( + respx_mock, monkeypatch, _rounds_logging_obj(recorder, "mcp-rounds-immediate") + ) + + await _settle(lambda: len(recorder.logged) > 0) + + _assert_one_row_bills_both_rounds(recorder.logged, response) + + +@pytest.mark.asyncio +async def test_non_stream_auto_execute_deferred_for_a_post_call_guardrail_logs_both_rounds_once( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + recorder: Final = _RoundRecorder() + logging_obj: Final = _rounds_logging_obj(recorder, "mcp-rounds-deferred") + logging_obj._defer_async_logging = True + response: Final = await _auto_execute_echo(respx_mock, monkeypatch, logging_obj) + + logging_obj._enqueue_deferred_logging() + await _settle(lambda: len(recorder.logged) > 0) + + _assert_one_row_bills_both_rounds(recorder.logged, response) + + +@pytest.mark.asyncio +async def test_non_stream_auto_execute_hands_sync_callbacks_the_client_response( + respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch +): + recorder: Final = _RoundRecorder() + response: Final = await _auto_execute_echo(respx_mock, monkeypatch, _rounds_logging_obj(recorder, "mcp-rounds-sync")) + + await _settle(lambda: len(recorder.hooked) > 0) + + assert len(recorder.hooked) == 1 + hooked: Final = recorder.hooked[0] + assert isinstance(hooked, ResponsesAPIResponse) + assert hooked.id == response.id + assert _output_types(hooked.output) == _output_types(response.output) + assert hooked.usage is not None + assert (hooked.usage.model_dump().get("prompt_tokens"), hooked.usage.model_dump().get("completion_tokens")) == ( + 71 + 149, + 19 + 5, + ) diff --git a/tests/unit/responses/mcp/test_round_billing.py b/tests/unit/responses/mcp/test_round_billing.py new file mode 100644 index 00000000000..cb73bace05e --- /dev/null +++ b/tests/unit/responses/mcp/test_round_billing.py @@ -0,0 +1,113 @@ +from types import MappingProxyType +from typing import Final + +import pytest + +from litellm.responses.mcp.round_billing import billed_for_every_round, merged_round_usage, summed_cost_breakdown +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse +from litellm.types.utils import CostBreakdown + + +def _round(response_id: str, response_cost: object) -> ResponsesAPIResponse: + response: Final = ResponsesAPIResponse.model_validate( + MappingProxyType( + { + "id": response_id, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.6", + "output": (), + "usage": MappingProxyType({"input_tokens": 10, "output_tokens": 2, "total_tokens": 12}), + } + ) + ) + response._hidden_params["response_cost"] = response_cost + return response + + +def test_summed_cost_breakdown_adds_each_charge_and_keeps_the_final_round_rates(): + first: Final[CostBreakdown] = { + "input_cost": 0.001, + "output_cost": 0.002, + "total_cost": 0.003, + "additional_costs": {"web_search": 0.01}, + "discount_percent": 0.1, + "margin_percent": 0.2, + "service_tier": "default", + } + final: Final[CostBreakdown] = { + "input_cost": 0.004, + "output_cost": 0.005, + "total_cost": 0.009, + "additional_costs": {"web_search": 0.01, "file_search": 0.0025}, + "discount_percent": 0.1, + "margin_percent": 0.2, + "service_tier": "default", + } + + summed: Final = summed_cost_breakdown(first, final) + + assert summed is not None + assert summed["input_cost"] == pytest.approx(0.005) + assert summed["output_cost"] == pytest.approx(0.007) + assert summed["total_cost"] == pytest.approx(0.012) + assert summed["additional_costs"] == {"web_search": pytest.approx(0.02), "file_search": pytest.approx(0.0025)} + assert (summed["discount_percent"], summed["margin_percent"]) == (0.1, 0.2) + assert summed["service_tier"] == "default" + + +def test_summed_cost_breakdown_is_none_when_the_final_round_was_not_repriced(): + breakdown: Final[CostBreakdown] = {"input_cost": 0.001, "output_cost": 0.002, "total_cost": 0.003} + + assert summed_cost_breakdown(breakdown, breakdown) is None + assert summed_cost_breakdown(None, breakdown) is None + + +def test_merged_round_usage_sums_token_details_and_keeps_provider_flags(): + first: Final = ResponseAPIUsage.model_validate( + MappingProxyType( + { + "input_tokens": 71, + "input_tokens_details": MappingProxyType({"cached_tokens": 32}), + "output_tokens": 19, + "total_tokens": 90, + "is_byok": True, + } + ) + ) + final: Final = ResponseAPIUsage.model_validate( + MappingProxyType( + { + "input_tokens": 149, + "input_tokens_details": MappingProxyType({"cached_tokens": 64}), + "output_tokens": 5, + "total_tokens": 154, + "is_byok": True, + } + ) + ) + + merged: Final = merged_round_usage(first, final) + + assert merged is not None + assert (merged.input_tokens, merged.output_tokens, merged.total_tokens) == (220, 24, 244) + assert merged.input_tokens_details is not None + assert merged.input_tokens_details.cached_tokens == 96 + assert merged.model_dump()["is_byok"] is True + + +def test_billed_for_every_round_leaves_the_cost_to_the_caller_when_a_round_has_no_price(): + billed: Final = billed_for_every_round(first=_round("resp_first", None), final=_round("resp_final", 0.004)) + + assert billed.id == "resp_final" + assert billed._hidden_params["response_cost"] is None + + +def test_billed_for_every_round_does_not_rebill_the_final_round_object(): + final: Final = _round("resp_final", 0.004) + + billed: Final = billed_for_every_round(first=_round("resp_first", 0.001), final=final) + + assert billed._hidden_params["response_cost"] == pytest.approx(0.005) + assert final._hidden_params["response_cost"] == 0.004