fix: isolate chained MCP spend logs

This commit is contained in:
tarunag10 2026-08-20 12:45:44 +05:30
parent 5290150a05
commit 03ca9fa932
4 changed files with 74 additions and 11 deletions

View file

@ -453,7 +453,9 @@ async def acompletion_with_mcp(
)
# Make follow-up call with streaming
follow_up_call_args: Final = dict(self.base_call_args)
follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(
self.base_call_args
)
follow_up_call_args["messages"] = follow_up_messages
follow_up_call_args["stream"] = True
# Ensure follow-up call doesn't trigger MCP handler again
@ -625,7 +627,9 @@ async def acompletion_with_mcp(
)
# Make follow-up call with original stream setting
follow_up_call_args: Final = dict(base_call_args)
follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(
base_call_args
)
follow_up_call_args["messages"] = follow_up_messages
follow_up_call_args["stream"] = stream

View file

@ -71,6 +71,28 @@ class LiteLLM_Proxy_MCP_Handler:
This handles when a user passes mcp server_url="litellm_proxy" in their tools.
"""
@staticmethod
def _prepare_follow_up_call_params(params: Mapping[str, Any]) -> dict[str, Any]:
"""Copy request params without state owned by the previous LLM call.
MCP auto-execution keeps the trace identifier so chained rounds remain
correlated, but each provider call must create its own logging object and
call identifier. Reusing either makes success dispatch one-shot and drops
spend rows for later rounds.
"""
follow_up_params = dict(params)
follow_up_params.pop("litellm_logging_obj", None)
follow_up_params.pop("litellm_call_id", None)
nested_params = follow_up_params.get("litellm_params")
if isinstance(nested_params, dict):
nested_params = dict(nested_params)
nested_params.pop("litellm_logging_obj", None)
nested_params.pop("litellm_call_id", None)
follow_up_params["litellm_params"] = nested_params
return follow_up_params
@staticmethod
def _get_parent_request_tags(kwargs: dict[str, Any] | None) -> list[str]:
"""Tags from the parent LLM request, using the same extraction logic as standard logging (incl. User-Agent)."""
@ -702,13 +724,19 @@ class LiteLLM_Proxy_MCP_Handler:
},
}
]
tool_logging_call_id = litellm_call_id or str(uuid.uuid4())
# SpendLogs.request_id is unique. The parent LLM call ID is
# therefore metadata, not the tool execution's call ID: reusing
# it causes all but the first MCP tool row to be skipped by the
# database's duplicate protection.
tool_logging_call_id = str(uuid.uuid4())
logging_metadata: dict[str, object] = {
"tool_call_id": tool_call_id,
"tool_name": sanitized_tool_name,
"server_name": server_name,
"headers": logging_safe_headers,
}
if litellm_call_id:
logging_metadata["parent_litellm_call_id"] = litellm_call_id
logging_request_data = {
"model": f"MCP: {tool_name}",
"metadata": logging_metadata,

View file

@ -782,7 +782,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
)
# Make follow-up call with streaming
follow_up_params: Final = self.original_request_params.copy()
follow_up_params: Final = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(
self.original_request_params
)
follow_up_params.update(
{
"input": follow_up_input,

View file

@ -410,7 +410,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
@pytest.mark.asyncio
async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_function_setup(
async def test_execute_tool_calls_uses_unique_call_ids_and_preserves_parent_context(
monkeypatch,
):
"""
@ -420,10 +420,10 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio
_setup_proxy_logging(monkeypatch)
call_tool_mock = _setup_mcp_call_environment(monkeypatch)
captured = {}
captured = []
def fake_function_setup(*_args, **kwargs):
captured.update(kwargs)
captured.append(kwargs)
return None, None
# NOTE: Don't patch via dotted string path here because `litellm.responses`
@ -435,7 +435,10 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
tool_name = "deepwiki-read_wiki_structure"
tool_calls = [{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}]
tool_calls = [
{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}},
{"id": "call-2", "function": {"name": tool_name, "arguments": "{}"}},
]
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map={tool_name: "deepwiki"},
@ -446,10 +449,36 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio
)
# Ensure the tool call was attempted (sanity)
assert call_tool_mock.await_count == 1
assert call_tool_mock.await_count == 2
assert len(captured) == 2
assert captured[0]["litellm_call_id"] != captured[1]["litellm_call_id"]
assert all(item["litellm_call_id"] != "cid" for item in captured)
assert all(item["litellm_trace_id"] == "tid" for item in captured)
assert all(item["metadata"]["parent_litellm_call_id"] == "cid" for item in captured)
assert captured.get("litellm_call_id") == "cid"
assert captured.get("litellm_trace_id") == "tid"
def test_prepare_follow_up_call_params_resets_per_call_logging_state():
original = {
"model": "gpt-4",
"litellm_call_id": "parent-call",
"litellm_logging_obj": object(),
"litellm_trace_id": "trace-1",
"litellm_params": {
"litellm_call_id": "nested-parent-call",
"litellm_logging_obj": object(),
"metadata": {"team": "legal"},
},
}
follow_up = LiteLLM_Proxy_MCP_Handler._prepare_follow_up_call_params(original)
assert "litellm_call_id" not in follow_up
assert "litellm_logging_obj" not in follow_up
assert "litellm_call_id" not in follow_up["litellm_params"]
assert "litellm_logging_obj" not in follow_up["litellm_params"]
assert follow_up["litellm_trace_id"] == "trace-1"
assert follow_up["litellm_params"]["metadata"] == {"team": "legal"}
assert original["litellm_call_id"] == "parent-call"
@pytest.mark.asyncio