From a873ead5d3c3d52e975bf2c6e8c2183b88cb7ae4 Mon Sep 17 00:00:00 2001 From: joshua Date: Sat, 19 Sep 2026 00:36:03 +0000 Subject: [PATCH] test(mcp): read SDK2 snake_case fields on CallToolResult Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../integrations/arize/test_arize_utils.py | 176 +++++------------- .../litellm_proxy/skills/test_skill_search.py | 4 +- .../test_cisco_ai_defense_mcp.py | 153 +++++---------- 3 files changed, 91 insertions(+), 242 deletions(-) diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py index 50f2823d632..165b7bc94d4 100644 --- a/tests/test_litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -70,9 +70,7 @@ def test_arize_set_attributes(): # Simulated LLM response object response_obj = ModelResponse( usage={"total_tokens": 100, "completion_tokens": 60, "prompt_tokens": 40}, - choices=[ - Choices(message={"role": "assistant", "content": "Basic Response Content"}) - ], + choices=[Choices(message={"role": "assistant", "content": "Basic Response Content"})], model="gpt-4o", id="chatcmpl-ID", ) @@ -89,9 +87,7 @@ def test_arize_set_attributes(): assert span.set_attribute.call_count == 26 # Metadata attached to the span - span.set_attribute.assert_any_call( - SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None}) - ) + span.set_attribute.assert_any_call(SpanAttributes.METADATA, json.dumps({"key_1": "value_1", "key_2": None})) # Basic LLM information span.set_attribute.assert_any_call(SpanAttributes.LLM_MODEL_NAME, "gpt-4o") @@ -114,16 +110,12 @@ def test_arize_set_attributes(): span.set_attribute.assert_any_call(SpanAttributes.OPENINFERENCE_SPAN_KIND, "LLM") # And TOOL must never be written for an LLM chat completion call. span_kind_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND ] assert "TOOL" not in span_kind_writes # Request message content and metadata - span.set_attribute.assert_any_call( - SpanAttributes.INPUT_VALUE, "Basic Request Content" - ) + span.set_attribute.assert_any_call(SpanAttributes.INPUT_VALUE, "Basic Request Content") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_INPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_ROLE}", "user", @@ -134,9 +126,7 @@ def test_arize_set_attributes(): ) # Tool call definitions and function names - span.set_attribute.assert_any_call( - f"{SpanAttributes.LLM_TOOLS}.0.name", "get_weather" - ) + span.set_attribute.assert_any_call(f"{SpanAttributes.LLM_TOOLS}.0.name", "get_weather") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_TOOLS}.0.description", "Fetches weather details.", @@ -146,26 +136,20 @@ def test_arize_set_attributes(): json.dumps( { "type": "object", - "properties": { - "location": {"type": "string", "description": "City name"} - }, + "properties": {"location": {"type": "string", "description": "City name"}}, "required": ["location"], } ), ) # Invocation parameters - span.set_attribute.assert_any_call( - SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}' - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_INVOCATION_PARAMETERS, '{"user": "test_user"}') # User ID span.set_attribute.assert_any_call(SpanAttributes.USER_ID, "test_user") # Output message content - span.set_attribute.assert_any_call( - SpanAttributes.OUTPUT_VALUE, "Basic Response Content" - ) + span.set_attribute.assert_any_call(SpanAttributes.OUTPUT_VALUE, "Basic Response Content") span.set_attribute.assert_any_call( f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_ROLE}", "assistant", @@ -228,9 +212,7 @@ def test_arize_set_attributes_responses_api(): ResponseReasoningItem( id="reasoning-001", type="reasoning", - summary=[ - Summary(text="First, I need to analyze...", type="summary_text") - ], + summary=[Summary(text="First, I need to analyze...", type="summary_text")], ), ResponseOutputMessage( id="msg-001", @@ -277,9 +259,7 @@ def test_arize_set_attributes_responses_api(): span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120) - span.set_attribute.assert_any_call( - SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180 - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180) def test_set_usage_outputs_pydantic_completion_usage(): @@ -327,9 +307,7 @@ def test_set_usage_outputs_pydantic_completion_usage(): span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 40) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 60) # reasoning_tokens for chat completions live in completion_tokens_details - span.set_attribute.assert_any_call( - SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25 - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25) def test_set_usage_outputs_pydantic_response_api_usage(): @@ -362,9 +340,7 @@ def test_set_usage_outputs_pydantic_response_api_usage(): span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120) span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250) - span.set_attribute.assert_any_call( - SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180 - ) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180) class TestArizeLogger(CustomLogger): @@ -375,16 +351,12 @@ class TestArizeLogger(CustomLogger): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.standard_callback_dynamic_params: Optional[ - StandardCallbackDynamicParams - ] = None + self.standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = None async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): # Capture dynamic params and print them for verification print("logged kwargs", json.dumps(kwargs, indent=4, default=str)) - self.standard_callback_dynamic_params = kwargs.get( - "standard_callback_dynamic_params" - ) + self.standard_callback_dynamic_params = kwargs.get("standard_callback_dynamic_params") @pytest.mark.asyncio @@ -410,14 +382,8 @@ async def test_arize_dynamic_params(): # Assert dynamic parameters were received in the callback assert test_arize_logger.standard_callback_dynamic_params is not None - assert ( - test_arize_logger.standard_callback_dynamic_params.get("arize_api_key") - == "test_api_key_dynamic" - ) - assert ( - test_arize_logger.standard_callback_dynamic_params.get("arize_space_key") - == "test_space_key_dynamic" - ) + assert test_arize_logger.standard_callback_dynamic_params.get("arize_api_key") == "test_api_key_dynamic" + assert test_arize_logger.standard_callback_dynamic_params.get("arize_space_key") == "test_space_key_dynamic" def test_construct_dynamic_arize_headers(): @@ -428,9 +394,7 @@ def test_construct_dynamic_arize_headers(): from litellm.types.utils import StandardCallbackDynamicParams # Test with all parameters present - dynamic_params_full = StandardCallbackDynamicParams( - arize_api_key="test_api_key", arize_space_id="test_space_id" - ) + dynamic_params_full = StandardCallbackDynamicParams(arize_api_key="test_api_key", arize_space_id="test_space_id") arize_logger = ArizeLogger() headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_full) @@ -438,9 +402,7 @@ def test_construct_dynamic_arize_headers(): assert headers == expected_headers # Test with only space_id - dynamic_params_space_id_only = StandardCallbackDynamicParams( - arize_space_id="test_space_id" - ) + dynamic_params_space_id_only = StandardCallbackDynamicParams(arize_space_id="test_space_id") headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_id_only) expected_headers = {"arize-space-id": "test_space_id"} @@ -456,9 +418,7 @@ def test_construct_dynamic_arize_headers(): dynamic_params_space_key_and_api_key = StandardCallbackDynamicParams( arize_space_key="test_space_key", arize_api_key="test_api_key" ) - headers = arize_logger.construct_dynamic_otel_headers( - dynamic_params_space_key_and_api_key - ) + headers = arize_logger.construct_dynamic_otel_headers(dynamic_params_space_key_and_api_key) expected_headers = {"arize-space-id": "test_space_key", "api_key": "test_api_key"} @@ -528,9 +488,7 @@ def test_arize_emits_no_cache_tokens_when_absent(): from litellm.integrations.arize._utils import _set_usage_outputs span = MagicMock() - response_obj = { - "usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6} - } + response_obj = {"usage": {"total_tokens": 10, "completion_tokens": 4, "prompt_tokens": 6}} _set_usage_outputs(span, response_obj, SpanAttributes) attrs = _collect_calls(span) assert SpanAttributes.LLM_TOKEN_COUNT_PROMPT_DETAILS_CACHE_READ not in attrs @@ -542,14 +500,8 @@ def test_passthrough_call_type_resolves_to_llm_span_kind(): from litellm.integrations._types.open_inference import OpenInferenceSpanKindValues from litellm.integrations.arize._utils import _infer_open_inference_span_kind - assert ( - _infer_open_inference_span_kind("allm_passthrough_route") - == OpenInferenceSpanKindValues.LLM.value - ) - assert ( - _infer_open_inference_span_kind("llm_passthrough_route") - == OpenInferenceSpanKindValues.LLM.value - ) + assert _infer_open_inference_span_kind("allm_passthrough_route") == OpenInferenceSpanKindValues.LLM.value + assert _infer_open_inference_span_kind("llm_passthrough_route") == OpenInferenceSpanKindValues.LLM.value def test_arize_chat_completion_with_tools_stays_llm_span_kind(): @@ -605,9 +557,7 @@ def test_arize_chat_completion_with_tools_stays_llm_span_kind(): ArizeLogger.set_arize_attributes(span, kwargs, response_obj) span_kind_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND ] assert span_kind_writes, "span.kind must be written" assert all(v == "LLM" for v in span_kind_writes) @@ -659,13 +609,8 @@ def test_arize_emits_assistant_tool_calls_on_output_message(): attrs = _collect_calls(span) base = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0.{MessageAttributes.MESSAGE_TOOL_CALLS}.0" assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc" - assert ( - attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather" - ) - assert ( - attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"] - == '{"location": "SF"}' - ) + assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_NAME}"] == "get_weather" + assert attrs[f"{base}.{ToolCallAttributes.TOOL_CALL_FUNCTION_ARGUMENTS_JSON}"] == '{"location": "SF"}' def test_arize_output_value_falls_back_to_tool_calls_summary(): @@ -818,9 +763,7 @@ def test_arize_emits_tool_call_id_and_name_on_input_tool_message(): assert attrs[f"{assistant_base}.{ToolCallAttributes.TOOL_CALL_ID}"] == "call_abc" # Tool message at index 2 tool_prefix = f"{SpanAttributes.LLM_INPUT_MESSAGES}.2" - assert ( - attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc" - ) + assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_TOOL_CALL_ID}"] == "call_abc" assert attrs[f"{tool_prefix}.{MessageAttributes.MESSAGE_NAME}"] == "get_weather" @@ -866,10 +809,7 @@ def test_arize_emits_multimodal_input_contents(): assert attrs[f"{base}.0.message_content.type"] == "text" assert attrs[f"{base}.0.message_content.text"] == "What is in this image?" assert attrs[f"{base}.1.message_content.type"] == "image" - assert ( - attrs[f"{base}.1.message_content.image.image.url"] - == "https://example.com/cat.png" - ) + assert attrs[f"{base}.1.message_content.image.image.url"] == "https://example.com/cat.png" def test_arize_emits_session_and_user_attrs_from_metadata(): @@ -974,11 +914,7 @@ def test_arize_does_not_overwrite_user_id_from_optional_params(): id="r2", ) ArizeLogger.set_arize_attributes(span, kwargs, response_obj) - user_id_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.USER_ID - ] + user_id_writes = [c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.USER_ID] assert "from_metadata" not in user_id_writes @@ -1048,9 +984,7 @@ def test_arize_passthrough_bedrock_anthropic_normalization(): "complete_input_dict": { "anthropic_version": "bedrock-2023-05-31", "max_tokens": 64, - "messages": [ - {"role": "user", "content": "What is the capital of France?"} - ], + "messages": [{"role": "user", "content": "What is the capital of France?"}], } }, "standard_logging_object": { @@ -1068,19 +1002,13 @@ def test_arize_passthrough_bedrock_anthropic_normalization(): assert attrs[SpanAttributes.INPUT_VALUE] == "What is the capital of France?" msg0 = f"{SpanAttributes.LLM_INPUT_MESSAGES}.0" assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_ROLE}"] == "user" - assert ( - attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"] - == "What is the capital of France?" - ) + assert attrs[f"{msg0}.{MessageAttributes.MESSAGE_CONTENT}"] == "What is the capital of France?" # Output rendering (Anthropic content[].text) assert attrs[SpanAttributes.OUTPUT_VALUE] == "The capital of France is Paris." out0 = f"{SpanAttributes.LLM_OUTPUT_MESSAGES}.0" assert attrs[f"{out0}.{MessageAttributes.MESSAGE_ROLE}"] == "assistant" - assert ( - attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"] - == "The capital of France is Paris." - ) + assert attrs[f"{out0}.{MessageAttributes.MESSAGE_CONTENT}"] == "The capital of France is Paris." # Token counts (Bedrock input_tokens/output_tokens) — extracted via # coercion of the non-dict response. @@ -1089,9 +1017,7 @@ def test_arize_passthrough_bedrock_anthropic_normalization(): # Span kind defended even though the call_type is a passthrough variant. span_kind_writes = [ - c.args[1] - for c in span.set_attribute.call_args_list - if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND + c.args[1] for c in span.set_attribute.call_args_list if c.args[0] == SpanAttributes.OPENINFERENCE_SPAN_KIND ] assert span_kind_writes # at least one assert all(v == "LLM" for v in span_kind_writes) @@ -1109,11 +1035,7 @@ def test_arize_passthrough_call_type_does_not_run_on_chat_completion(): span = MagicMock() _maybe_normalize_passthrough( span, - { - "additional_args": { - "complete_input_dict": {"messages": [{"role": "user", "content": "x"}]} - } - }, + {"additional_args": {"complete_input_dict": {"messages": [{"role": "user", "content": "x"}]}}}, {"choices": [{"message": {"role": "assistant", "content": "y"}}]}, {"choices": [{"message": {"role": "assistant", "content": "y"}}]}, {"call_type": "completion"}, @@ -1133,11 +1055,7 @@ def test_arize_passthrough_skipped_when_message_redaction_enabled(): span = MagicMock() kwargs = { "additional_args": { - "complete_input_dict": { - "messages": [ - {"role": "user", "content": "Patient John Doe, SSN 123-45-6789"} - ] - } + "complete_input_dict": {"messages": [{"role": "user", "content": "Patient John Doe, SSN 123-45-6789"}]} }, # Enables redaction via the dynamic-param path inside # should_redact_message_logging(), without touching globals. @@ -1211,9 +1129,7 @@ def test_arize_mcp_call_tool_result_does_not_break_attribute_setting(): "optional_params": {}, "litellm_params": {"custom_llm_provider": "mcp"}, } - response_obj = CallToolResult( - content=[TextContent(type="text", text="sunny, 21C")], isError=False - ) + response_obj = CallToolResult(content=[TextContent(type="text", text="sunny, 21C")], is_error=False) ArizeLogger.set_arize_attributes(span, kwargs, response_obj) @@ -1231,11 +1147,11 @@ def test_arize_coerce_response_obj_dumps_pydantic_without_get(): from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs - result = CallToolResult(content=[TextContent(type="text", text="hi")], isError=False) + result = CallToolResult(content=[TextContent(type="text", text="hi")], is_error=False) coerced = _coerce_response_obj_for_attrs(result) assert isinstance(coerced, dict) - assert coerced["isError"] is False + assert coerced["is_error"] is False assert coerced["content"][0]["text"] == "hi" @@ -1295,9 +1211,7 @@ def test_arize_mcp_tool_span_renders_name_input_and_output(): from mcp.types import CallToolResult, TextContent span = MagicMock() - response_obj = CallToolResult( - content=[TextContent(type="text", text="sunny, 21C")], isError=False - ) + response_obj = CallToolResult(content=[TextContent(type="text", text="sunny, 21C")], is_error=False) ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) @@ -1318,7 +1232,7 @@ def test_arize_mcp_tool_span_serializes_non_text_content(): span = MagicMock() response_obj = CallToolResult( content=[ImageContent(type="image", data="Zm9v", mimeType="image/png")], - isError=False, + is_error=False, ) ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) @@ -1336,9 +1250,7 @@ def test_arize_mcp_tool_span_respects_message_redaction(): from mcp.types import CallToolResult, TextContent span = MagicMock() - response_obj = CallToolResult( - content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False - ) + response_obj = CallToolResult(content=[TextContent(type="text", text="SSN 123-45-6789")], is_error=False) ArizeLogger.set_arize_attributes( span, @@ -1390,7 +1302,7 @@ def test_arize_mcp_tool_span_renders_empty_arguments(): span = MagicMock() kwargs = _mcp_kwargs(mcp_tool_call_metadata={"name": "ping", "arguments": {}}) - response_obj = CallToolResult(content=[TextContent(type="text", text="pong")], isError=False) + response_obj = CallToolResult(content=[TextContent(type="text", text="pong")], is_error=False) ArizeLogger.set_arize_attributes(span, kwargs, response_obj) @@ -1405,7 +1317,7 @@ def test_arize_mcp_tool_span_renders_empty_content(): from mcp.types import CallToolResult span = MagicMock() - response_obj = CallToolResult(content=[], isError=False) + response_obj = CallToolResult(content=[], is_error=False) ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) @@ -1420,7 +1332,7 @@ def test_arize_mcp_tool_span_falls_back_to_structured_content(): from mcp.types import CallToolResult span = MagicMock() - response_obj = CallToolResult(content=[], structuredContent={"temp_c": 21}, isError=False) + response_obj = CallToolResult(content=[], structured_content={"temp_c": 21}, is_error=False) ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) @@ -1463,7 +1375,7 @@ def test_arize_mcp_tool_span_serializes_mixed_text_and_media(): TextContent(type="text", text="see image"), ImageContent(type="image", data="Zm9v", mimeType="image/png"), ], - isError=False, + is_error=False, ) ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) diff --git a/tests/test_litellm/llms/litellm_proxy/skills/test_skill_search.py b/tests/test_litellm/llms/litellm_proxy/skills/test_skill_search.py index a0f22a59f0c..3f1fe0d5d68 100644 --- a/tests/test_litellm/llms/litellm_proxy/skills/test_skill_search.py +++ b/tests/test_litellm/llms/litellm_proxy/skills/test_skill_search.py @@ -420,7 +420,7 @@ class TestHandleSkillSearchMCP: result = await handle_skill_search( query="language translation", top_k=10_000, user_api_key_dict=UserAPIKeyAuth(user_id="u") ) - assert result.isError is False + assert result.is_error is False assert len(json.loads(result.content[0].text)) == MAX_SKILL_SEARCH_TOP_K @pytest.mark.asyncio @@ -432,5 +432,5 @@ class TestHandleSkillSearchMCP: result = await handle_skill_search( query="language translation", top_k=0, user_api_key_dict=UserAPIKeyAuth(user_id="u") ) - assert result.isError is False + assert result.is_error is False assert len(json.loads(result.content[0].text)) == 1 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py index 137b7d24023..07436199a8d 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cisco_ai_defense_mcp.py @@ -51,9 +51,7 @@ class TestCiscoAIDefenseMCPMode: @pytest.mark.asyncio async def test_mcp_mode_inspects_mcp_request(self): g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call") - data = _mcp_request( - name="send_email", args={"to": "x@y.com"}, litellm_call_id="call-1" - ) + data = _mcp_request(name="send_email", args={"to": "x@y.com"}, litellm_call_id="call-1") post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL)) with _patch_inspection_post(g, post_mock): result = await g.async_pre_call_hook( @@ -78,9 +76,7 @@ class TestCiscoAIDefenseMCPMode: async def test_mcp_mode_blocks_violation(self): g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call") data = _mcp_request(name="leak_secrets", args={"target": "evil"}) - with _patch_inspection_post( - g, AsyncMock(return_value=_violation_response(url=MCP_URL)) - ): + with _patch_inspection_post(g, AsyncMock(return_value=_violation_response(url=MCP_URL))): with pytest.raises(HTTPException) as exc: await g.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth(), @@ -165,9 +161,7 @@ class TestCiscoAIDefenseMCPMode: call_type="mcp_call", ) - forwarded = ProxyLogging( - user_api_key_cache=UserApiKeyCache() - )._convert_mcp_hook_response_to_kwargs( + forwarded = ProxyLogging(user_api_key_cache=UserApiKeyCache())._convert_mcp_hook_response_to_kwargs( response_data=result, original_kwargs={"arguments": dict(original_args)} ) assert forwarded["arguments"] == sanitized_args, ( @@ -179,14 +173,10 @@ class TestCiscoAIDefenseMCPMode: @pytest.mark.asyncio async def test_mcp_response_hook_inspects_tool_output(self): - g = _make_guardrail( - inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"] - ) + g = _make_guardrail(inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]) response_obj = _mcp_response( - SimpleNamespace( - content=[{"type": "text", "text": "Here is the secret API key abc123"}] - ) + SimpleNamespace(content=[{"type": "text", "text": "Here is the secret API key abc123"}]) ) post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL)) @@ -215,9 +205,7 @@ class TestCiscoAIDefenseMCPMode: "name": "lookup_secret", "arguments": {"key": "production"}, } - assert sent_payload["result"]["content"][0]["text"] == ( - "Here is the secret API key abc123" - ) + assert sent_payload["result"]["content"][0]["text"] == ("Here is the secret API key abc123") assert "request" not in sent_payload assert "metadata" not in sent_payload @@ -225,12 +213,8 @@ class TestCiscoAIDefenseMCPMode: async def test_mcp_response_hook_blocks_violation(self): from litellm.types.mcp import MCPPostCallResponseObject - g = _make_guardrail( - inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"] - ) - response_obj = _mcp_response( - SimpleNamespace(content=[{"type": "text", "text": "leaked"}]) - ) + g = _make_guardrail(inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]) + response_obj = _mcp_response(SimpleNamespace(content=[{"type": "text", "text": "leaked"}])) post_mock = AsyncMock(return_value=_violation_response(url=MCP_URL)) with _patch_inspection_post(g, post_mock): @@ -257,9 +241,7 @@ class TestCiscoAIDefenseMCPMode: @pytest.mark.asyncio async def test_mcp_response_hook_skipped_in_chat_mode(self): g = _make_guardrail() - response_obj = _mcp_response( - SimpleNamespace(content=[{"type": "text", "text": "hi"}]) - ) + response_obj = _mcp_response(SimpleNamespace(content=[{"type": "text", "text": "hi"}])) post_mock = AsyncMock() with _patch_inspection_post(g, post_mock): @@ -291,11 +273,7 @@ class TestCiscoAIDefenseMCPMode: @pytest.mark.asyncio async def test_mcp_response_hook_runs_with_pre_mcp_call_only(self): g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call") - response_obj = _mcp_response( - SimpleNamespace( - content=[{"type": "text", "text": "would have been scanned"}] - ) - ) + response_obj = _mcp_response(SimpleNamespace(content=[{"type": "text", "text": "would have been scanned"}])) post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL)) with _patch_inspection_post(g, post_mock): @@ -317,26 +295,18 @@ class TestCiscoAIDefenseMCPMode: [("safe", False), ("violation", True)], ) @pytest.mark.asyncio - async def test_mcp_response_hook_handles_raw_list_content( - self, cisco_response_kind, expected_block - ): + async def test_mcp_response_hook_handles_raw_list_content(self, cisco_response_kind, expected_block): from litellm.types.mcp import MCPPostCallResponseObject - g = _make_guardrail( - inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"] - ) + g = _make_guardrail(inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]) text_content = ( - "exfiltrated data: ..." - if cisco_response_kind == "violation" - else "Here is the secret API key abc123" + "exfiltrated data: ..." if cisco_response_kind == "violation" else "Here is the secret API key abc123" ) response_obj = _mcp_response([{"type": "text", "text": text_content}]) cisco_resp = ( - _violation_response(url=MCP_URL) - if cisco_response_kind == "violation" - else _safe_response(url=MCP_URL) + _violation_response(url=MCP_URL) if cisco_response_kind == "violation" else _safe_response(url=MCP_URL) ) post_mock = AsyncMock(return_value=cisco_resp) kwargs = { @@ -354,8 +324,7 @@ class TestCiscoAIDefenseMCPMode: ) assert post_mock.called, ( - "MCP response inspect was silently skipped for raw-list " - "shape — _normalize_mcp_response failed." + "MCP response inspect was silently skipped for raw-list shape — _normalize_mcp_response failed." ) assert post_mock.call_args.kwargs["url"] == MCP_URL @@ -382,14 +351,12 @@ class TestCiscoAIDefenseMCPMode: from litellm.types.mcp import MCPPostCallResponseObject - g = _make_guardrail( - inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"] - ) + g = _make_guardrail(inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]) real_result = CallToolResult( content=[TextContent(type="text", text="leak 9045629876")], - structuredContent={"patient": {"ssn": "123-45-6789"}}, - isError=False, + structured_content={"patient": {"ssn": "123-45-6789"}}, + is_error=False, ) wrapped = MCPPostCallResponseObject( mcp_tool_call_response=real_result, @@ -397,12 +364,8 @@ class TestCiscoAIDefenseMCPMode: ) assert isinstance(wrapped.mcp_tool_call_response, list) - assert all( - isinstance(item, tuple) and len(item) == 2 - for item in wrapped.mcp_tool_call_response - ), ( - "Pydantic coercion shape changed — update the normalizer to " - "match the new wire format." + assert all(isinstance(item, tuple) and len(item) == 2 for item in wrapped.mcp_tool_call_response), ( + "Pydantic coercion shape changed — update the normalizer to match the new wire format." ) post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL)) @@ -441,9 +404,7 @@ class TestCiscoAIDefenseMCPMode: f"``content`` field." ) assert content_items[0].get("type") == "text" - assert sent_payload["result"]["structuredContent"] == { - "patient": {"ssn": "123-45-6789"} - } + assert sent_payload["result"]["structuredContent"] == {"patient": {"ssn": "123-45-6789"}} assert sent_payload["result"]["isError"] is False assert sent_payload["id"] == "real-wire-call" assert sent_payload["method"] == "tools/call" @@ -482,7 +443,6 @@ class TestCiscoAIDefenseMCPMode: class TestCiscoAIDefenseRedactListShape: - @staticmethod def _violation_with_redact_response(text: str = "[REDACTED tool output]"): return _mock_inspect_response( @@ -512,8 +472,8 @@ class TestCiscoAIDefenseRedactListShape: tuples_list = [ ("meta", None), ("content", inner_content), - ("structuredContent", {"patient": {"ssn": "123-45-6789"}}), - ("isError", False), + ("structured_content", {"patient": {"ssn": "123-45-6789"}}), + ("is_error", False), ] return tuples_list, lambda: inner_content[0].text @@ -526,16 +486,12 @@ class TestCiscoAIDefenseRedactListShape: from litellm.types.mcp import MCPPostCallResponseObject - g = _make_guardrail( - inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"] - ) + g = _make_guardrail(inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]) content, get_text = getattr(self, factory_name)() response_obj = _mcp_response(content) - with _patch_inspection_post( - g, AsyncMock(return_value=self._violation_with_redact_response()) - ): + with _patch_inspection_post(g, AsyncMock(return_value=self._violation_with_redact_response())): result = await g.async_post_mcp_tool_call_hook( kwargs={"name": "leak", "arguments": {}}, response_obj=response_obj, @@ -544,15 +500,13 @@ class TestCiscoAIDefenseRedactListShape: ) assert result is None or not isinstance(result, MCPPostCallResponseObject), ( - f"Redact silently fell through to block for {factory_name}. " - f"result={result!r}" + f"Redact silently fell through to block for {factory_name}. result={result!r}" ) assert get_text() == "[REDACTED tool output]", ( - f"Redact silently failed for {factory_name}; original text " - f"not rewritten." + f"Redact silently failed for {factory_name}; original text not rewritten." ) if factory_name == "_pydantic_tuple_list_factory": - structured_content = dict(content)["structuredContent"] + structured_content = dict(content)["structured_content"] assert structured_content == {"result": "[REDACTED tool output]"} assert "123-45-6789" not in json.dumps(structured_content) @@ -565,20 +519,16 @@ class TestCiscoAIDefenseRedactListShape: original_response = CallToolResult( content=[TextContent(type="text", text="SSN: 123-45-6789")], - structuredContent={"patient": {"ssn": "123-45-6789"}}, - isError=False, + structured_content={"patient": {"ssn": "123-45-6789"}}, + is_error=False, ) wrapper = MCPPostCallResponseObject( mcp_tool_call_response=original_response, hidden_params=HiddenParams(), ) - g = _make_guardrail( - inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"] - ) - with _patch_inspection_post( - g, AsyncMock(return_value=self._violation_with_redact_response()) - ): + g = _make_guardrail(inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]) + with _patch_inspection_post(g, AsyncMock(return_value=self._violation_with_redact_response())): await g.async_post_mcp_tool_call_hook( kwargs={ "name": "leak", @@ -591,12 +541,12 @@ class TestCiscoAIDefenseRedactListShape: ) assert original_response.content[0].text == "[REDACTED tool output]" - assert "123-45-6789" not in json.dumps(original_response.structuredContent), ( + assert "123-45-6789" not in json.dumps(original_response.structured_content), ( "Redact verdict left the client-visible MCP tool output unchanged. " "The post-call hook receives a wrapped MCPPostCallResponseObject but " "the endpoint returns kwargs['original_response'], so the redaction " "must rewrite that object too. structuredContent still leaks: " - f"{original_response.structuredContent!r}" + f"{original_response.structured_content!r}" ) @@ -606,9 +556,7 @@ class TestCiscoAIDefenseMcpInputRedactionFallback: @pytest.mark.asyncio async def test_single_string_arg_is_rewritten(self): g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call") - data = _mcp_request( - name="search", args={"query": "my SSN is 123-45-6789", "limit": 10} - ) + data = _mcp_request(name="search", args={"query": "my SSN is 123-45-6789", "limit": 10}) cisco = _redact_response(sanitized_text="my SSN is [REDACTED]", url=MCP_URL) with _patch_inspection_post(g, AsyncMock(return_value=cisco)): result = await g.async_pre_call_hook( @@ -663,7 +611,6 @@ class TestCiscoAIDefenseMcpInputRedactionFallback: class TestCiscoAIDefenseMCPBlockingContract: - @pytest.mark.asyncio async def test_block_response_survives_dispatcher_contract(self): from litellm.litellm_core_utils.litellm_logging import Logging @@ -677,8 +624,8 @@ class TestCiscoAIDefenseMCPBlockingContract: ) raw_response = CallToolResult( content=[TextContent(type="text", text="exfiltrated")], - structuredContent={"result": "exfiltrated"}, - isError=False, + structured_content={"result": "exfiltrated"}, + is_error=False, ) response_obj = MCPPostCallResponseObject( mcp_tool_call_response=raw_response, @@ -712,11 +659,11 @@ class TestCiscoAIDefenseMCPBlockingContract: "Hook must keep returning a MCPPostCallResponseObject for " "dispatcher paths that do honor returned replacements." ) - assert raw_response.isError is True + assert raw_response.is_error is True assert "Blocked by Cisco AI Defense" in raw_response.content[0].text - assert raw_response.structuredContent is not None - assert "Blocked by Cisco AI Defense" in raw_response.structuredContent["result"] - assert "exfiltrated" not in raw_response.structuredContent["result"] + assert raw_response.structured_content is not None + assert "Blocked by Cisco AI Defense" in raw_response.structured_content["result"] + assert "exfiltrated" not in raw_response.structured_content["result"] logging_stub = Logging.__new__(Logging) logging_stub.model_call_details = {} parsed = logging_stub._parse_post_mcp_call_hook_response(response=result) @@ -725,7 +672,6 @@ class TestCiscoAIDefenseMCPBlockingContract: class TestCiscoAIDefenseJsonRpcSuccessEnvelope: - @staticmethod def _cisco_mcp_envelope(*, is_safe: bool, action: str = "Block") -> Response: return _mock_inspect_response( @@ -761,12 +707,8 @@ class TestCiscoAIDefenseJsonRpcSuccessEnvelope: ], ) @pytest.mark.asyncio - async def test_mcp_jsonrpc_envelope_respects_verdict( - self, is_safe, action, should_block - ): - g = _make_guardrail( - name="cisco-mcp", inspection_type="mcp", event_hook="pre_mcp_call" - ) + async def test_mcp_jsonrpc_envelope_respects_verdict(self, is_safe, action, should_block): + g = _make_guardrail(name="cisco-mcp", inspection_type="mcp", event_hook="pre_mcp_call") data = _mcp_request( name="ask_question", args={ @@ -776,9 +718,7 @@ class TestCiscoAIDefenseJsonRpcSuccessEnvelope: ) with _patch_inspection_post( g, - AsyncMock( - return_value=self._cisco_mcp_envelope(is_safe=is_safe, action=action) - ), + AsyncMock(return_value=self._cisco_mcp_envelope(is_safe=is_safe, action=action)), ): if should_block: with pytest.raises(HTTPException) as exc: @@ -790,10 +730,7 @@ class TestCiscoAIDefenseJsonRpcSuccessEnvelope: ) assert exc.value.status_code == 400 assert exc.value.detail["surface"] == "mcp" - assert ( - exc.value.detail["event_id"] - == "645d9d22-b016-47e0-a12c-9d587fb11c57" - ) + assert exc.value.detail["event_id"] == "645d9d22-b016-47e0-a12c-9d587fb11c57" else: result = await g.async_pre_call_hook( user_api_key_dict=UserAPIKeyAuth(),