mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
4a7d8bbffa
commit
a873ead5d3
3 changed files with 91 additions and 242 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue