This commit is contained in:
devin-ai-integration[bot] 2026-09-30 12:13:02 -07:00 • committed by GitHub
commit ff5598e533
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 490 additions and 54 deletions

View file

@ -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."""

View file

@ -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,

View file

@ -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

View 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

View file

@ -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,

View file

@ -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,
)

View 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