test(mcp): read SDK2 snake_case fields on CallToolResult

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-19 00:36:03 +00:00
parent 4a7d8bbffa
commit a873ead5d3
3 changed files with 91 additions and 242 deletions

View file

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

View file

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

View file

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