mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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
This commit is contained in:
parent
85dc7cb62e
commit
350b88d9ff
7 changed files with 490 additions and 54 deletions
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
83
litellm/responses/mcp/round_billing.py
Normal file
83
litellm/responses/mcp/round_billing.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
113
tests/unit/responses/mcp/test_round_billing.py
Normal file
113
tests/unit/responses/mcp/test_round_billing.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue