mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix: mark MCP tool isError responses as failures
Signed-off-by: Ritwij Aryan Parmar <ritwij.aryan.parmar@gmail.com>
This commit is contained in:
parent
5699a06413
commit
3e033d1eac
2 changed files with 194 additions and 1 deletions
|
|
@ -5368,6 +5368,58 @@ def _extract_response_obj_and_hidden_params(
|
|||
return response_obj, hidden_params
|
||||
|
||||
|
||||
def _extract_mcp_tool_error_message(response_obj: dict) -> Optional[str]:
|
||||
"""Return an upstream MCP tool error message when result.isError is true."""
|
||||
result_obj = response_obj.get("result", response_obj)
|
||||
if not isinstance(result_obj, dict) or result_obj.get("isError") is not True:
|
||||
return None
|
||||
|
||||
content = result_obj.get("content")
|
||||
messages: List[str] = []
|
||||
if isinstance(content, list):
|
||||
for item in content:
|
||||
if isinstance(item, dict):
|
||||
text = item.get("text") or item.get("message")
|
||||
if isinstance(text, str) and text.strip():
|
||||
messages.append(text.strip())
|
||||
elif isinstance(item, str) and item.strip():
|
||||
messages.append(item.strip())
|
||||
elif isinstance(content, str) and content.strip():
|
||||
messages.append(content.strip())
|
||||
|
||||
if messages:
|
||||
return " ".join(messages)
|
||||
return "MCP tool returned isError=true"
|
||||
|
||||
|
||||
def _get_mcp_tool_call_logging_status(
|
||||
call_type: Optional[str],
|
||||
response_obj: dict,
|
||||
status: StandardLoggingPayloadStatus,
|
||||
error_str: Optional[str],
|
||||
) -> Tuple[
|
||||
Optional[str],
|
||||
StandardLoggingPayloadStatus,
|
||||
Optional[str],
|
||||
Optional[StandardLoggingPayloadErrorInformation],
|
||||
]:
|
||||
if call_type != CallTypes.call_mcp_tool.value:
|
||||
return call_type, status, error_str, None
|
||||
|
||||
mcp_tool_error_message = _extract_mcp_tool_error_message(response_obj)
|
||||
if mcp_tool_error_message is None:
|
||||
return call_type, status, error_str, None
|
||||
|
||||
error_information = StandardLoggingPayloadErrorInformation(
|
||||
error_code="",
|
||||
error_class="MCPToolError",
|
||||
llm_provider="mcp",
|
||||
traceback="",
|
||||
error_message=mcp_tool_error_message,
|
||||
)
|
||||
return call_type, "failure", error_str or mcp_tool_error_message, error_information
|
||||
|
||||
|
||||
def get_standard_logging_object_payload(
|
||||
kwargs: Optional[dict],
|
||||
init_response_obj: Union[Any, BaseModel, dict],
|
||||
|
|
@ -5396,7 +5448,14 @@ def get_standard_logging_object_payload(
|
|||
)
|
||||
|
||||
completion_start_time = kwargs.get("completion_start_time", end_time)
|
||||
call_type = kwargs.get("call_type")
|
||||
call_type, status, error_str, mcp_tool_error_information = (
|
||||
_get_mcp_tool_call_logging_status(
|
||||
call_type=kwargs.get("call_type"),
|
||||
response_obj=response_obj,
|
||||
status=status,
|
||||
error_str=error_str,
|
||||
)
|
||||
)
|
||||
cache_hit = kwargs.get("cache_hit", False)
|
||||
# Extract usage as a plain dict, avoiding Pydantic round-trip
|
||||
usage_dict = StandardLoggingPayloadSetup.get_usage_as_dict(
|
||||
|
|
@ -5488,6 +5547,8 @@ def get_standard_logging_object_payload(
|
|||
error_information = StandardLoggingPayloadSetup.get_error_information(
|
||||
original_exception=original_exception,
|
||||
)
|
||||
if mcp_tool_error_information is not None:
|
||||
error_information = mcp_tool_error_information
|
||||
|
||||
## get final response object ##
|
||||
final_response_obj = StandardLoggingPayloadSetup.get_final_response_obj(
|
||||
|
|
|
|||
|
|
@ -2689,6 +2689,138 @@ def test_get_standard_logging_object_payload_includes_litellm_call_id(logging_ob
|
|||
assert payload["litellm_call_id"] == call_id
|
||||
|
||||
|
||||
def test_mcp_tool_is_error_marks_standard_logging_payload_as_failure(logging_obj):
|
||||
import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
now = datetime.datetime.now()
|
||||
payload = get_standard_logging_object_payload(
|
||||
kwargs={
|
||||
"litellm_call_id": "mcp-call-id",
|
||||
"model": "MCP: deepwiki/search",
|
||||
"messages": [],
|
||||
"call_type": CallTypes.call_mcp_tool.value,
|
||||
"mcp_tool_call_metadata": {
|
||||
"name": "search",
|
||||
"arguments": {"NOT_a_real_arg": "x"},
|
||||
"mcp_server_name": "deepwiki",
|
||||
"namespaced_tool_name": "deepwiki/search",
|
||||
},
|
||||
},
|
||||
init_response_obj={
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Input validation error: 'query' is required",
|
||||
}
|
||||
],
|
||||
"isError": True,
|
||||
},
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
logging_obj=logging_obj,
|
||||
status="success",
|
||||
)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["status_fields"]["llm_api_status"] == "failure"
|
||||
assert payload["error_str"] == "Input validation error: 'query' is required"
|
||||
assert payload["error_information"] is not None
|
||||
assert payload["error_information"]["error_class"] == "MCPToolError"
|
||||
assert (
|
||||
payload["error_information"]["error_message"]
|
||||
== "Input validation error: 'query' is required"
|
||||
)
|
||||
assert payload["metadata"]["mcp_tool_call_metadata"] is not None
|
||||
assert (
|
||||
payload["metadata"]["mcp_tool_call_metadata"]["namespaced_tool_name"]
|
||||
== "deepwiki/search"
|
||||
)
|
||||
|
||||
|
||||
def test_mcp_tool_json_rpc_is_error_envelope_marks_payload_as_failure(logging_obj):
|
||||
import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
now = datetime.datetime.now()
|
||||
payload = get_standard_logging_object_payload(
|
||||
kwargs={
|
||||
"litellm_call_id": "mcp-json-rpc-call-id",
|
||||
"model": "MCP: github/create_issue",
|
||||
"messages": [],
|
||||
"call_type": CallTypes.call_mcp_tool.value,
|
||||
},
|
||||
init_response_obj={
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"result": {
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Permission denied for repository",
|
||||
}
|
||||
],
|
||||
"isError": True,
|
||||
},
|
||||
},
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
logging_obj=logging_obj,
|
||||
status="success",
|
||||
)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["error_information"] is not None
|
||||
assert payload["error_information"]["error_class"] == "MCPToolError"
|
||||
assert (
|
||||
payload["error_information"]["error_message"]
|
||||
== "Permission denied for repository"
|
||||
)
|
||||
|
||||
|
||||
def test_mcp_tool_success_keeps_standard_logging_payload_success(logging_obj):
|
||||
import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
now = datetime.datetime.now()
|
||||
payload = get_standard_logging_object_payload(
|
||||
kwargs={
|
||||
"litellm_call_id": "mcp-success-call-id",
|
||||
"model": "MCP: deepwiki/search",
|
||||
"messages": [],
|
||||
"call_type": CallTypes.call_mcp_tool.value,
|
||||
},
|
||||
init_response_obj={
|
||||
"content": [{"type": "text", "text": "Search result"}],
|
||||
"isError": False,
|
||||
},
|
||||
start_time=now,
|
||||
end_time=now,
|
||||
logging_obj=logging_obj,
|
||||
status="success",
|
||||
)
|
||||
|
||||
assert payload is not None
|
||||
assert payload["status"] == "success"
|
||||
assert payload["error_str"] is None
|
||||
assert payload["error_information"] is not None
|
||||
assert payload["error_information"]["error_message"] == ""
|
||||
|
||||
|
||||
def _make_dict_logging_obj():
|
||||
"""Build a Logging instance configured for a non-streaming dict result."""
|
||||
obj = LitellmLogging(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue