diff --git a/enterprise/enterprise_hooks/aporia_ai.py b/enterprise/enterprise_hooks/aporia_ai.py index d2184e92f2f..de741aa6ca7 100644 --- a/enterprise/enterprise_hooks/aporia_ai.py +++ b/enterprise/enterprise_hooks/aporia_ai.py @@ -173,6 +173,7 @@ class AporiaGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): from litellm.proxy.common_utils.callback_utils import ( diff --git a/enterprise/enterprise_hooks/google_text_moderation.py b/enterprise/enterprise_hooks/google_text_moderation.py index fe26a03207f..61987af7532 100644 --- a/enterprise/enterprise_hooks/google_text_moderation.py +++ b/enterprise/enterprise_hooks/google_text_moderation.py @@ -95,6 +95,7 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): """ diff --git a/enterprise/enterprise_hooks/openai_moderation.py b/enterprise/enterprise_hooks/openai_moderation.py index ee8ac495099..0b6f34018b4 100644 --- a/enterprise/enterprise_hooks/openai_moderation.py +++ b/enterprise/enterprise_hooks/openai_moderation.py @@ -42,6 +42,7 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): text = "" diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py b/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py index a44af55d4b1..ea428b51b8e 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/llama_guard.py @@ -105,6 +105,7 @@ class _ENTERPRISE_LlamaGuard(CustomLogger): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): """ diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py b/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py index 1475a94303e..6735998960b 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/llm_guard.py @@ -127,6 +127,7 @@ class _ENTERPRISE_LLMGuard(CustomLogger): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): """ diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py index 00230937b32..1028a443a42 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/pagerduty/pagerduty.py @@ -147,6 +147,7 @@ class PagerDutyAlerting(SlackAlerting): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Union[Exception, str, dict]]: """ diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index a2c788a8f21..e069a89b9c5 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -290,6 +290,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): "aretrieve_fine_tuning_job", "alist_fine_tuning_jobs", "acancel_fine_tuning_job", + "mcp_call", ], ) -> Union[Exception, str, Dict, None]: """ diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index b6792354334..501185b207e 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -234,7 +234,6 @@ class CustomGuardrail(CustomLogger): Returns True if the guardrail should be run on the event_type """ requested_guardrails = self.get_guardrail_from_metadata(data) - verbose_logger.debug( "inside should_run_guardrail for guardrail=%s event_type= %s guardrail_supported_event_hooks= %s requested_guardrails= %s self.default_on= %s", self.guardrail_name, @@ -243,7 +242,6 @@ class CustomGuardrail(CustomLogger): requested_guardrails, self.default_on, ) - if self.default_on is True: if self._event_hook_is_event_type(event_type): if isinstance(self.event_hook, Mode): @@ -287,7 +285,6 @@ class CustomGuardrail(CustomLogger): ) if result is not None: return result - return True def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool: diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index cdc12005471..ded5ccca766 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -281,6 +281,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[ Union[Exception, str, dict] @@ -327,6 +328,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac "moderation", "audio_transcription", "responses", + "mcp_call", ], ) -> Any: pass diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 3c5ccb40515..0abd0232d21 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -17,6 +17,8 @@ from mcp.types import CallToolResult from mcp.types import Tool as MCPTool from litellm._logging import verbose_logger +from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException +from fastapi import HTTPException from litellm.experimental_mcp_client.client import MCPClient from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, @@ -592,22 +594,22 @@ class MCPServerManager: "server_name": server_name_from_prefix, "user_api_key_auth": user_api_key_auth, } - pre_hook_result = await proxy_logging_obj.async_pre_mcp_tool_call_hook( - kwargs=pre_hook_kwargs, - request_obj=None, # Will be created in the hook - start_time=start_time, - end_time=start_time, - ) - - if pre_hook_result: - # Check if the call should proceed - if not pre_hook_result.get("should_proceed", True): - error_message = pre_hook_result.get("error_message", "Tool call rejected by pre-hook") - raise ValueError(error_message) + try: + pre_hook_result = await proxy_logging_obj.async_pre_mcp_tool_call_hook( + kwargs=pre_hook_kwargs, + request_obj=None, # Will be created in the hook + start_time=start_time, + end_time=start_time, + ) - # Apply any argument modifications - if pre_hook_result.get("modified_arguments"): - arguments = pre_hook_result["modified_arguments"] + if pre_hook_result: + # Apply any argument modifications + if pre_hook_result.get("modified_arguments"): + arguments = pre_hook_result["modified_arguments"] + except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e: + # Re-raise guardrail exceptions to properly fail the MCP call + verbose_logger.error(f"Guardrail blocked MCP tool call pre call: {str(e)}") + raise e # Get server-specific auth header if available server_auth_header = None @@ -627,6 +629,7 @@ class MCPServerManager: ) async with client: + # Use the original tool name (without prefix) for the actual call call_tool_params = MCPCallToolRequestParams( name=original_tool_name, @@ -635,40 +638,39 @@ class MCPServerManager: # Initialize during_hook_task as None during_hook_task = None - + tasks = [] # Start during hook if proxy_logging_obj is available if proxy_logging_obj: - try: - during_hook_task = asyncio.create_task( - proxy_logging_obj.async_during_mcp_tool_call_hook( - kwargs={ - "name": name, - "arguments": arguments, - "server_name": server_name_from_prefix, - }, - request_obj=None, # Will be created in the hook - start_time=start_time, - end_time=start_time, - ) + during_hook_task = asyncio.create_task( + proxy_logging_obj.async_during_mcp_tool_call_hook( + kwargs={ + "name": name, + "arguments": arguments, + "server_name": server_name_from_prefix, + }, + request_obj=None, # Will be created in the hook + start_time=start_time, + end_time=start_time, ) - except Exception as e: - verbose_logger.warning(f"During hook error (non-blocking): {str(e)}") + ) + tasks.append(during_hook_task) - result = await client.call_tool(call_tool_params) - - ######################################################### - # Check during hook result if it completed - ######################################################### - if proxy_logging_obj and during_hook_task is not None: - try: - during_hook_result = await during_hook_task - if during_hook_result and not during_hook_result.get("should_continue", True): - error_message = during_hook_result.get("error_message", "Tool call cancelled by during-hook") - raise ValueError(error_message) - except Exception as e: - verbose_logger.warning(f"During hook error (non-blocking): {str(e)}") - - return result + + tasks.append(asyncio.create_task(client.call_tool(call_tool_params))) + try: + + mcp_responses = await asyncio.gather(*tasks) + + # If proxy_logging_obj is None, the tool call result is at index 0 + # If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task) + result_index = 1 if proxy_logging_obj else 0 + result = mcp_responses[result_index] + + return cast(CallToolResult, result) + except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e: + # Re-raise guardrail exceptions to properly fail the MCP call + verbose_logger.error(f"Guardrail blocked MCP tool call during result check: {str(e)}") + raise e ######################################################### # End of Methods that call the upstream MCP servers diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index c4783c6df00..a31fabf57cf 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -141,15 +141,52 @@ if MCP_AVAILABLE: REST API to call a specific MCP tool with the provided arguments """ from litellm.proxy.proxy_server import add_litellm_data_to_request, proxy_config + from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException + from fastapi import HTTPException - data = await request.json() - data = await add_litellm_data_to_request( - data=data, - request=request, - user_api_key_dict=user_api_key_dict, - proxy_config=proxy_config, - ) - return await call_mcp_tool(**data) + try: + data = await request.json() + data = await add_litellm_data_to_request( + data=data, + request=request, + user_api_key_dict=user_api_key_dict, + proxy_config=proxy_config, + ) + return await call_mcp_tool(**data) + except BlockedPiiEntityError as e: + verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}") + raise HTTPException( + status_code=400, + detail={ + "error": "blocked_pii_entity", + "message": str(e), + "entity_type": getattr(e, 'entity_type', None), + "guardrail_name": getattr(e, 'guardrail_name', None) + } + ) + except GuardrailRaisedException as e: + verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") + raise HTTPException( + status_code=400, + detail={ + "error": "guardrail_violation", + "message": str(e), + "guardrail_name": getattr(e, 'guardrail_name', None) + } + ) + except HTTPException as e: + # Re-raise HTTPException as-is to preserve status code and detail + verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") + raise e + except Exception as e: + verbose_logger.exception(f"Unexpected error in MCP tool call: {str(e)}") + raise HTTPException( + status_code=500, + detail={ + "error": "internal_server_error", + "message": f"An unexpected error occurred: {str(e)}" + } + ) ######################################################## # MCP Connection testing routes diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 62c365f92c4..dd4ac9b007d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -218,6 +218,7 @@ if MCP_AVAILABLE: from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request from litellm.proxy.proxy_server import proxy_config + from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException # Validate arguments user_api_key_auth, mcp_auth_header, _, mcp_server_auth_headers, mcp_protocol_version = get_auth_context() @@ -254,9 +255,34 @@ if MCP_AVAILABLE: mcp_protocol_version=mcp_protocol_version, **data, # for logging ) + except BlockedPiiEntityError as e: + verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}") + # Return error as text content for MCP protocol + return [TextContent( + text=f"Error: Blocked PII entity detected - {str(e)}", + type="text" + )] + except GuardrailRaisedException as e: + verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") + # Return error as text content for MCP protocol + return [TextContent( + text=f"Error: Guardrail violation - {str(e)}", + type="text" + )] + except HTTPException as e: + verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") + # Return error as text content for MCP protocol + return [TextContent( + text=f"Error: {str(e.detail)}", + type="text" + )] except Exception as e: verbose_logger.exception(f"MCP mcp_server_tool_call - error: {e}") - raise e + # Return error as text content for MCP protocol + return [TextContent( + text=f"Error: {str(e)}", + type="text" + )] return response diff --git a/litellm/proxy/example_config_yaml/custom_callbacks1.py b/litellm/proxy/example_config_yaml/custom_callbacks1.py index 83f68dd55a2..d84980f42d7 100644 --- a/litellm/proxy/example_config_yaml/custom_callbacks1.py +++ b/litellm/proxy/example_config_yaml/custom_callbacks1.py @@ -30,6 +30,7 @@ class MyCustomHandler( "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ): return data @@ -63,6 +64,7 @@ class MyCustomHandler( "moderation", "audio_transcription", "responses", + "mcp_call", ], ): pass diff --git a/litellm/proxy/example_config_yaml/custom_guardrail.py b/litellm/proxy/example_config_yaml/custom_guardrail.py index 5a5c7844107..b05498eae5f 100644 --- a/litellm/proxy/example_config_yaml/custom_guardrail.py +++ b/litellm/proxy/example_config_yaml/custom_guardrail.py @@ -32,6 +32,7 @@ class myCustomGuardrail(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Union[Exception, str, dict]]: """ @@ -67,6 +68,7 @@ class myCustomGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 317db76e37d..621defa2d4c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -76,6 +76,7 @@ class AimGuardrail(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Union[Exception, str, dict, None]: verbose_proxy_logger.debug("Inside AIM Pre-Call Hook") @@ -95,6 +96,7 @@ class AimGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ) -> Union[Exception, str, dict, None]: verbose_proxy_logger.debug("Inside AIM Moderation Hook") diff --git a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py index 85a0ab9dedd..0d41e40e833 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py @@ -188,6 +188,7 @@ class AporiaGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): from litellm.proxy.common_utils.callback_utils import ( diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index fc589263f57..c0bdc06dc60 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -123,6 +123,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Dict[str, Any]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py index 9c31e062028..4c6f8d335a7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/text_moderation.py @@ -218,6 +218,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Dict[str, Any]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index b287232d49f..384d958946b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -73,6 +73,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): """ + # Set supported event hooks to include MCP hooks + if 'supported_event_hooks' not in kwargs: + kwargs['supported_event_hooks'] = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + GuardrailEventHooks.during_call, + GuardrailEventHooks.pre_mcp_call, + GuardrailEventHooks.during_mcp_call, + ] + super().__init__(**kwargs) BaseAWSLLM.__init__(self) @@ -400,6 +410,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Union[Exception, str, dict, None]: verbose_proxy_logger.debug("Inside AIM Pre-Call Hook") @@ -458,6 +469,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): from litellm.proxy.common_utils.callback_utils import ( diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py index 87860477f0b..91878ad0f0d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_guardrail.py @@ -35,6 +35,7 @@ class myCustomGuardrail(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Union[Exception, str, dict]]: """ @@ -71,6 +72,7 @@ class myCustomGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py index ea5b7641c6b..db22f630adc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/guardrails_ai/guardrails_ai.py @@ -204,6 +204,7 @@ class GuardrailsAI(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[ Union[Exception, str, dict] diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py index 2dd8a3154a8..bd0a9eba7d5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py @@ -135,6 +135,7 @@ class lakeraAI_Moderation(CustomGuardrail): "pass_through_endpoint", "rerank", "responses", + "mcp_call", ], ): if ( @@ -313,6 +314,7 @@ class lakeraAI_Moderation(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Union[Exception, str, Dict]]: from litellm.types.guardrails import GuardrailEventHooks @@ -347,6 +349,7 @@ class lakeraAI_Moderation(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): if self.event_hook is None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py index e7d7d3b5aaa..33c1526d1c8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai_v2.py @@ -190,6 +190,7 @@ class LakeraAIGuardrail(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Union[Exception, str, Dict]]: from litellm.proxy.common_utils.callback_utils import ( @@ -261,6 +262,7 @@ class LakeraAIGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): from litellm.proxy.common_utils.callback_utils import ( diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index 00350d5a36f..cddbdc010d5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -82,6 +82,7 @@ class LassoGuardrail(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Union[Exception, str, dict, None]: verbose_proxy_logger.debug("Inside Lasso Pre-Call Hook") @@ -99,6 +100,7 @@ class LassoGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ): """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 889d02bb736..ee04e899c6d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -225,6 +225,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Union[Exception, str, dict, None]: """Pre-call hook to sanitize user prompts.""" diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py index 53e5c7a472b..fb2552471cb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py @@ -193,6 +193,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Dict[str, Any]]: """ @@ -247,6 +248,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ) -> Optional[Dict[str, Any]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 41b3c1368a5..bc3a6e1a7cb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -245,6 +245,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Dict[str, Any]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py index f4df62f70af..f4741aa8e00 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py +++ b/litellm/proxy/guardrails/guardrail_hooks/pillar/pillar.py @@ -142,6 +142,7 @@ class PillarGuardrail(CustomGuardrail): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[Union[Exception, str, dict]]: """ @@ -188,6 +189,7 @@ class PillarGuardrail(CustomGuardrail): "moderation", "audio_transcription", "responses", + "mcp_call", ], ) -> Optional[Union[Exception, str, dict]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index d6b958513ba..84778740203 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -403,11 +403,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if call_type in [ LitellmCallTypes.completion.value, LitellmCallTypes.acompletion.value, - ]: - + ] or call_type == "mcp_call": messages = data["messages"] tasks = [] - for m in messages: content = m.get("content", None) if content is None: diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index e06366d02b5..6c72b4c9d24 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -197,6 +197,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): "audio_transcription", "pass_through_endpoint", "rerank", + "mcp_call", ], ) -> Optional[ Union[Exception, str, dict] diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 4486b99d2ec..51c5014563b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -55,7 +55,7 @@ from litellm import ( from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes from litellm.caching.caching import DualCache, RedisCache -from litellm.exceptions import RejectedRequestError +from litellm.exceptions import RejectedRequestError, BlockedPiiEntityError, GuardrailRaisedException from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting @@ -457,9 +457,10 @@ class ProxyLogging: Pre MCP Tool Call Hook Use this to validate and modify MCP tool calls before execution. + Reuses existing LLM guardrail logic by converting MCP calls to message format. """ from litellm.types.llms.base import HiddenParams - from litellm.types.mcp import MCPPreCallRequestObject, MCPPreCallResponseObject + from litellm.types.mcp import MCPPreCallRequestObject callbacks = self.get_combined_callback_list( dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None), @@ -468,33 +469,64 @@ class ProxyLogging: # Create the request object if it's not already one if not isinstance(request_obj, MCPPreCallRequestObject): + # Convert UserAPIKeyAuth object to dict if needed + user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth")) + request_obj = MCPPreCallRequestObject( tool_name=kwargs.get("name", ""), arguments=kwargs.get("arguments", {}), server_name=kwargs.get("server_name"), - user_api_key_auth=kwargs.get("user_api_key_auth"), + user_api_key_auth=user_api_key_auth_dict, hidden_params=HiddenParams(), ) for callback in callbacks: try: - if isinstance(callback, CustomLogger): - response: Optional[MCPPreCallResponseObject] = ( - await callback.async_pre_mcp_tool_call_hook( - kwargs=kwargs, - request_obj=request_obj, - start_time=start_time, - end_time=end_time, - ) + _callback: Optional[CustomLogger] = None + if isinstance(callback, str): + from typing import cast + from litellm import _custom_logger_compatible_callbacks_literal + _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( + cast(_custom_logger_compatible_callbacks_literal, callback) ) - ###################################################################### - # if any of the callbacks return a response, use the first one - # this allows for validation failures or argument modifications - ###################################################################### - if response is not None: - return self._parse_pre_mcp_call_hook_response( - response=response, original_request=request_obj + else: + _callback = callback # type: ignore + + if _callback is not None and isinstance(_callback, CustomGuardrail): + from litellm.types.guardrails import GuardrailEventHooks + + # Check if guardrail should be run for pre_call hook (reusing existing logic) + if ( + _callback.should_run_guardrail( + data=kwargs, event_type=GuardrailEventHooks.pre_mcp_call ) + is not True + ): + continue + + # Convert MCP tool call to LLM message format for existing guardrail logic + synthetic_llm_data = self._convert_mcp_to_llm_format(request_obj, kwargs) + # Reuse existing LLM guardrail logic + user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth")) + + result = await _callback.async_pre_call_hook( + user_api_key_dict=user_api_key_auth_dict, + cache=self.call_details["user_api_key_cache"], + data=synthetic_llm_data, + call_type="mcp_call" + ) + + # Convert result back to MCP response format if blocked/modified + if result is not None: + mcp_response = self._convert_llm_result_to_mcp_response(result, request_obj) + if mcp_response is not None: + return self._parse_pre_mcp_call_hook_response( + response=mcp_response, original_request=request_obj + ) + + except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e: + # Re-raise guardrail exceptions so they can be properly handled + raise e except Exception as e: verbose_proxy_logger.exception( "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format( @@ -503,6 +535,224 @@ class ProxyLogging: ) return None + def _convert_user_api_key_auth_to_dict(self, user_api_key_auth_obj): + """ + Helper function to convert UserAPIKeyAuth object to dictionary. + Handles both Pydantic models and regular objects. + """ + if user_api_key_auth_obj is not None: + if hasattr(user_api_key_auth_obj, 'model_dump'): + # If it's a Pydantic model, convert to dict + return user_api_key_auth_obj.model_dump() + elif hasattr(user_api_key_auth_obj, '__dict__'): + # If it's a regular object, convert to dict + return user_api_key_auth_obj.__dict__ + return user_api_key_auth_obj + + def _convert_mcp_to_llm_format(self, request_obj, kwargs: dict) -> dict: + """ + Convert MCP tool call to LLM message format for existing guardrail validation. + """ + from litellm.types.llms.openai import ChatCompletionUserMessage + + # Create a synthetic message that represents the tool call + tool_call_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" + + synthetic_message = ChatCompletionUserMessage( + role="user", + content=tool_call_content + ) + + # Create synthetic LLM data that guardrails can process + synthetic_data = { + "messages": [synthetic_message], + "model": kwargs.get("model", "mcp-tool-call"), + "user_api_key_user_id": kwargs.get("user_api_key_user_id"), + "user_api_key_team_id": kwargs.get("user_api_key_team_id"), + "user_api_key_end_user_id": kwargs.get("user_api_key_end_user_id"), + "user_api_key_hash": kwargs.get("user_api_key_hash"), + "user_api_key_request_route": kwargs.get("user_api_key_request_route"), + "mcp_tool_name": request_obj.tool_name, # Keep original for reference + "mcp_arguments": request_obj.arguments, # Keep original for reference + } + + return synthetic_data + + def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> Optional[Any]: + """ + Convert LLM guardrail result back to MCP response format. + """ + from litellm.types.mcp import MCPPreCallResponseObject + + # If result is an exception, it means the guardrail blocked the request + if isinstance(llm_result, Exception): + return MCPPreCallResponseObject( + should_proceed=False, + error_message=str(llm_result), + modified_arguments=None + ) + + # If result is a dict with modified messages, check for content filtering + if isinstance(llm_result, dict): + modified_messages = llm_result.get("messages") + if modified_messages: + # Check if content was blocked/modified + original_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" + new_content = modified_messages[0].get("content", "") if modified_messages else "" + + if new_content != original_content: + # Content was modified - could be masking, redaction, or blocking + if not new_content or "blocked" in new_content.lower() or "violation" in new_content.lower(): + # Content was blocked completely + return MCPPreCallResponseObject( + should_proceed=False, + error_message="Content blocked by guardrail", + modified_arguments=None + ) + else: + # Content was masked/redacted - extract the modified arguments + try: + # Try to parse the modified arguments from the masked content + modified_args = self._extract_modified_arguments_from_content(new_content, request_obj) + if modified_args is not None: + # Return the masked/redacted arguments for the MCP call to use + return MCPPreCallResponseObject( + should_proceed=True, + error_message=None, + modified_arguments=modified_args + ) + else: + # Could not parse modified arguments, allow original call but warn + verbose_proxy_logger.warning( + f"Could not parse modified arguments from guardrail response: {new_content}" + ) + return None + except Exception as e: + verbose_proxy_logger.error(f"Error parsing modified arguments: {e}") + # Fallback: allow original call + return None + + # If result is a string, it's likely an error message + if isinstance(llm_result, str): + return MCPPreCallResponseObject( + should_proceed=False, + error_message=llm_result, + modified_arguments=None + ) + + return None + + def _extract_modified_arguments_from_content(self, masked_content: str, request_obj) -> Optional[dict]: + """ + Extract modified/masked arguments from the guardrail response content. + """ + import json + + verbose_proxy_logger.debug(f"Extracting modified args from content: {masked_content}") + + try: + # The format should be: "Tool: \nArguments: " + # Parse the arguments section + lines = masked_content.strip().split('\n') + for i, line in enumerate(lines): + if line.startswith("Arguments:"): + # Get the arguments part - everything after "Arguments: " + args_text = line[len("Arguments:"):].strip() + + verbose_proxy_logger.debug(f"Found arguments text: {args_text}") + + # Try to parse as JSON first + try: + modified_args = json.loads(args_text) + verbose_proxy_logger.debug(f"Successfully parsed JSON args: {modified_args}") + return modified_args + except json.JSONDecodeError as e: + # If JSON parsing fails, try to extract key-value pairs manually + verbose_proxy_logger.debug(f"Failed to parse JSON arguments: {args_text}, error: {e}") + return self._parse_arguments_manually(args_text, request_obj.arguments) + + # If we can't find the Arguments: line, return None + verbose_proxy_logger.warning("Could not find 'Arguments:' line in masked content") + return None + + except Exception as e: + verbose_proxy_logger.error(f"Error extracting modified arguments: {e}") + return None + + def _parse_arguments_manually(self, args_text: str, original_args: dict) -> Optional[dict]: + """ + Try to manually parse arguments when JSON parsing fails. + This is a fallback for cases where the guardrail modifies the format. + """ + import re + + try: + # Start with original arguments and try to apply modifications + modified_args = original_args.copy() + + # Look for simple key-value patterns + # This is a basic implementation - can be enhanced based on specific guardrail formats + for key, original_value in original_args.items(): + if isinstance(original_value, str): + # Look for the key in the masked content and try to extract its value + pattern = rf"['\"]?{re.escape(key)}['\"]?\s*:\s*['\"]?([^,'\"]*)['\"]?" + match = re.search(pattern, args_text, re.IGNORECASE) + if match: + new_value = match.group(1).strip() + if new_value: + modified_args[key] = new_value + + return modified_args + + except Exception as e: + verbose_proxy_logger.error(f"Error in manual argument parsing: {e}") + return None + + def _convert_llm_result_to_mcp_during_response(self, llm_result, request_obj) -> Optional[Any]: + """ + Convert LLM guardrail result back to MCP during call response format. + """ + from litellm.types.mcp import MCPDuringCallResponseObject + + # If result is an exception, it means the guardrail wants to stop execution + if isinstance(llm_result, Exception): + return MCPDuringCallResponseObject( + should_continue=False, + error_message=str(llm_result) + ) + + # If result is a dict with modified messages, check for content filtering + if isinstance(llm_result, dict): + modified_messages = llm_result.get("messages") + if modified_messages: + # Check if content was blocked/modified + original_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" + new_content = modified_messages[0].get("content", "") if modified_messages else "" + + if new_content != original_content: + # Content was modified, could be masking or blocking + if not new_content or "blocked" in new_content.lower(): + # Content was blocked + return MCPDuringCallResponseObject( + should_continue=False, + error_message="Content blocked by guardrail during execution" + ) + else: + # Content was masked/modified - for now, stop execution + return MCPDuringCallResponseObject( + should_continue=False, + error_message="Content modified by guardrail during execution" + ) + + # If result is a string, it's likely an error message + if isinstance(llm_result, str): + return MCPDuringCallResponseObject( + should_continue=False, + error_message=llm_result + ) + + return None + def get_combined_callback_list( self, dynamic_success_callbacks: Optional[List], global_callbacks: List ) -> List: @@ -542,12 +792,11 @@ class ProxyLogging: During MCP Tool Call Hook Use this for concurrent monitoring and validation during tool execution. + Reuses existing LLM guardrail logic by converting MCP calls to message format. """ from litellm.types.llms.base import HiddenParams - from litellm.types.mcp import ( - MCPDuringCallRequestObject, - MCPDuringCallResponseObject, - ) + from litellm.types.mcp import MCPDuringCallRequestObject + callbacks = self.get_combined_callback_list( dynamic_success_callbacks=getattr(self, "dynamic_success_callbacks", None), @@ -566,24 +815,47 @@ class ProxyLogging: for callback in callbacks: try: - if isinstance(callback, CustomLogger): - response: Optional[MCPDuringCallResponseObject] = ( - await callback.async_during_mcp_tool_call_hook( - kwargs=kwargs, - request_obj=request_obj, - start_time=start_time, - end_time=end_time, - ) + _callback: Optional[CustomLogger] = None + if isinstance(callback, str): + from typing import cast + from litellm import _custom_logger_compatible_callbacks_literal + _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( + cast(_custom_logger_compatible_callbacks_literal, callback) ) - ###################################################################### - # if any of the callbacks return a response, use the first one - # this allows for execution control decisions - ###################################################################### - if response is not None: - return self._parse_during_mcp_call_hook_response( - response=response + else: + _callback = callback # type: ignore + + if _callback is not None and isinstance(_callback, CustomGuardrail): + from litellm.types.guardrails import GuardrailEventHooks + + # Check if guardrail should be run for during_call hook (reusing existing logic) + if ( + _callback.should_run_guardrail( + data=kwargs, event_type=GuardrailEventHooks.during_mcp_call ) + is not True + ): + continue + # Convert MCP tool call to LLM message format for existing guardrail logic + synthetic_llm_data = self._convert_mcp_to_llm_format(request_obj, kwargs) + + # Reuse existing LLM guardrail logic for during call + user_api_key_auth_dict = self._convert_user_api_key_auth_to_dict(kwargs.get("user_api_key_auth")) + + result = await _callback.async_moderation_hook( + data=synthetic_llm_data, + user_api_key_dict=user_api_key_auth_dict, + call_type="mcp_call" + ) + # Convert result back to MCP response format if blocked/modified + if result is not None: + mcp_response = self._convert_llm_result_to_mcp_during_response(result, request_obj) + if mcp_response is not None: + return self._parse_during_mcp_call_hook_response(response=mcp_response) + + except Exception as e: + raise e verbose_proxy_logger.exception( "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {}".format( str(e) diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 65a7099e038..5c72b9b6521 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -181,6 +181,8 @@ class LiteLLM_Proxy_MCP_Handler: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) + from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException + from fastapi import HTTPException tool_results = [] tool_call_id: Optional[str] = None @@ -211,6 +213,27 @@ class LiteLLM_Proxy_MCP_Handler: "result": result_text }) + except BlockedPiiEntityError as e: + verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}") + error_message = f"Tool call blocked: PII entity '{getattr(e, 'entity_type', 'unknown')}' detected by guardrail '{getattr(e, 'guardrail_name', 'unknown')}'. {str(e)}" + tool_results.append({ + "tool_call_id": tool_call_id, + "result": error_message + }) + except GuardrailRaisedException as e: + verbose_logger.error(f"GuardrailRaisedException in MCP tool call: {str(e)}") + error_message = f"Tool call blocked: Guardrail '{getattr(e, 'guardrail_name', 'unknown')}' violation. {str(e)}" + tool_results.append({ + "tool_call_id": tool_call_id, + "result": error_message + }) + except HTTPException as e: + verbose_logger.error(f"HTTPException in MCP tool call: {str(e)}") + error_message = f"Tool call failed: {str(e.detail) if hasattr(e, 'detail') else str(e)}" + tool_results.append({ + "tool_call_id": tool_call_id, + "result": error_message + }) except Exception as e: verbose_logger.exception(f"Error executing MCP tool call: {e}") tool_results.append({ diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index bb9a5ff8894..fd18484a898 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -491,6 +491,8 @@ class GuardrailEventHooks(str, Enum): post_call = "post_call" during_call = "during_call" logging_only = "logging_only" + pre_mcp_call = "pre_mcp_call" + during_mcp_call = "during_mcp_call" class DynamicGuardrailParams(TypedDict): diff --git a/tests/mcp_tests/test_mcp_guardrails.py b/tests/mcp_tests/test_mcp_guardrails.py new file mode 100644 index 00000000000..83febcf7dcb --- /dev/null +++ b/tests/mcp_tests/test_mcp_guardrails.py @@ -0,0 +1,734 @@ +""" +Test file for MCP Guardrails Feature + +This file tests the MCP guardrails functionality for both pre and during MCP call hooks, +including various guardrail types and proper exception handling. +""" + +import asyncio +import pytest +import sys +import os +from datetime import datetime +from typing import Optional, Dict, Any +from unittest.mock import MagicMock, AsyncMock, patch + +# Add the project root to the path +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.caching.caching import DualCache +from litellm.types.mcp import ( + MCPPreCallRequestObject, + MCPPreCallResponseObject, + MCPDuringCallRequestObject, + MCPDuringCallResponseObject, +) +from litellm.types.llms.base import HiddenParams +from litellm.types.guardrails import GuardrailEventHooks +from fastapi import HTTPException + + +class MockPiiGuardrail(CustomGuardrail): + """Mock PII guardrail that raises BlockedPiiEntityError""" + + def __init__(self, should_block: bool = True, entity_type: str = "EMAIL_ADDRESS"): + super().__init__() + self.should_block = should_block + self.entity_type = entity_type + self.guardrail_name = "mock-pii-guardrail" + self.call_count = 0 + + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + """Always run for testing""" + return True + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ): + """Mock pre-call hook that raises BlockedPiiEntityError""" + self.call_count += 1 + + if self.should_block: + raise BlockedPiiEntityError( + entity_type=self.entity_type, + guardrail_name=self.guardrail_name, + ) + return None + + +class MockContentGuardrail(CustomGuardrail): + """Mock content guardrail that raises GuardrailRaisedException""" + + def __init__(self, should_block: bool = True): + super().__init__() + self.should_block = should_block + self.guardrail_name = "mock-content-guardrail" + self.call_count = 0 + + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + """Always run for testing""" + return True + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ): + """Mock pre-call hook that raises GuardrailRaisedException""" + self.call_count += 1 + + if self.should_block: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="Content violates policy" + ) + return None + + +class MockHttpGuardrail(CustomGuardrail): + """Mock HTTP guardrail that raises HTTPException""" + + def __init__(self, should_block: bool = True): + super().__init__() + self.should_block = should_block + self.guardrail_name = "mock-http-guardrail" + self.call_count = 0 + + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + """Always run for testing""" + return True + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ): + """Mock pre-call hook that raises HTTPException""" + self.call_count += 1 + + if self.should_block: + raise HTTPException( + status_code=400, + detail={"error": "Violated guardrail policy"} + ) + return None + + +class MockDuringCallGuardrail(CustomGuardrail): + """Mock guardrail for during-call testing""" + + def __init__(self, should_block: bool = True): + super().__init__() + self.should_block = should_block + self.guardrail_name = "mock-during-guardrail" + self.call_count = 0 + + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + """Always run for testing""" + return True + + async def async_moderation_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + call_type: str, + ): + """Mock during-call hook that raises exceptions""" + self.call_count += 1 + + if self.should_block: + raise BlockedPiiEntityError( + entity_type="PHONE_NUMBER", + guardrail_name=self.guardrail_name, + ) + return None + + +class MockProxyLogging: + """Mock proxy logging object for testing MCP guardrails""" + + def __init__(self, guardrails: Optional[list] = None): + self.guardrails = guardrails if guardrails is not None else [] + self.call_details = {"user_api_key_cache": DualCache()} + self.dynamic_success_callbacks = [] + self.call_count = 0 + + def get_combined_callback_list(self, dynamic_success_callbacks, global_callbacks): + """Return the guardrails for testing""" + return self.guardrails + + def _convert_mcp_to_llm_format(self, request_obj, kwargs: dict) -> dict: + """Convert MCP tool call to LLM message format""" + tool_call_content = f"Tool: {request_obj.tool_name}\nArguments: {request_obj.arguments}" + + return { + "messages": [{"role": "user", "content": tool_call_content}], + "model": kwargs.get("model", "mcp-tool-call"), + "user_api_key_user_id": kwargs.get("user_api_key_user_id"), + "user_api_key_team_id": kwargs.get("user_api_key_team_id"), + } + + def _convert_llm_result_to_mcp_response(self, llm_result, request_obj): + """Convert LLM result back to MCP response format""" + return None # For testing, we don't need to convert back + + def _parse_pre_mcp_call_hook_response(self, response, original_request): + """Parse pre MCP call hook response""" + return response + + async def async_pre_mcp_tool_call_hook( + self, + kwargs: dict, + request_obj: Any, + start_time: datetime, + end_time: datetime, + ) -> Optional[Any]: + """Mock pre MCP tool call hook""" + self.call_count += 1 + + # Simulate the actual hook logic + for guardrail in self.guardrails: + if isinstance(guardrail, CustomGuardrail): + try: + synthetic_data = self._convert_mcp_to_llm_format(request_obj, kwargs) + + # Check if guardrail should run + if not guardrail.should_run_guardrail(synthetic_data, GuardrailEventHooks.pre_mcp_call): + continue + + result = await guardrail.async_pre_call_hook( + user_api_key_dict=kwargs.get("user_api_key_auth"), + cache=self.call_details["user_api_key_cache"], + data=synthetic_data, + call_type="mcp_call" + ) + if result is not None: + return self._parse_pre_mcp_call_hook_response(result, request_obj) + except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e: + # Re-raise guardrail exceptions + raise e + except Exception as e: + # Log non-guardrail exceptions as non-blocking + print(f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {str(e)}") + + return None + + async def async_during_mcp_tool_call_hook( + self, + kwargs: dict, + request_obj: Any, + start_time: datetime, + end_time: datetime, + ) -> Optional[Any]: + """Mock during MCP tool call hook""" + self.call_count += 1 + + # Simulate the actual hook logic + for guardrail in self.guardrails: + if isinstance(guardrail, CustomGuardrail): + try: + synthetic_data = self._convert_mcp_to_llm_format(request_obj, kwargs) + result = await guardrail.async_moderation_hook( + data=synthetic_data, + user_api_key_dict=kwargs.get("user_api_key_auth"), + call_type="mcp_call" + ) + if result is not None: + return result + except (BlockedPiiEntityError, GuardrailRaisedException, HTTPException) as e: + # Re-raise guardrail exceptions + raise e + except Exception as e: + # Log non-guardrail exceptions as non-blocking + print(f"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging {str(e)}") + + return None + + +@pytest.fixture +def mock_user_api_key(): + """Mock user API key for testing""" + return UserAPIKeyAuth(api_key="test_key", user_id="test_user") + + +@pytest.fixture +def mock_cache(): + """Mock cache for testing""" + return DualCache() + + +@pytest.fixture +def mock_pii_guardrail(): + """Mock PII guardrail that blocks""" + return MockPiiGuardrail(should_block=True) + + +@pytest.fixture +def mock_pii_guardrail_allow(): + """Mock PII guardrail that allows""" + return MockPiiGuardrail(should_block=False) + + +@pytest.fixture +def mock_content_guardrail(): + """Mock content guardrail that blocks""" + return MockContentGuardrail(should_block=True) + + +@pytest.fixture +def mock_http_guardrail(): + """Mock HTTP guardrail that blocks""" + return MockHttpGuardrail(should_block=True) + + +@pytest.fixture +def mock_during_guardrail(): + """Mock during-call guardrail that blocks""" + return MockDuringCallGuardrail(should_block=True) + + +@pytest.fixture +def mock_proxy_logging(): + """Mock proxy logging object""" + return MockProxyLogging() + + +class TestMCPGuardrailsPreCall: + """Test MCP guardrails for pre-call hooks""" + + @pytest.mark.asyncio + async def test_pii_guardrail_blocks_pre_call(self, mock_pii_guardrail, mock_user_api_key, mock_cache): + """Test that PII guardrail properly blocks pre-call""" + proxy_logging = MockProxyLogging([mock_pii_guardrail]) + + # Create MCP request + request_obj = MCPPreCallRequestObject( + tool_name="email_tool", + arguments={"email": "test@example.com"}, + server_name="email_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "email_tool", + "arguments": {"email": "test@example.com"}, + "server_name": "email_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that BlockedPiiEntityError is raised + with pytest.raises(BlockedPiiEntityError) as excinfo: + await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Verify the error details + assert excinfo.value.entity_type == "EMAIL_ADDRESS" + assert excinfo.value.guardrail_name == "mock-pii-guardrail" + assert mock_pii_guardrail.call_count == 1 + + @pytest.mark.asyncio + async def test_pii_guardrail_allows_pre_call(self, mock_pii_guardrail_allow, mock_user_api_key, mock_cache): + """Test that PII guardrail allows pre-call when configured to allow""" + proxy_logging = MockProxyLogging([mock_pii_guardrail_allow]) + + request_obj = MCPPreCallRequestObject( + tool_name="email_tool", + arguments={"email": "test@example.com"}, + server_name="email_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "email_tool", + "arguments": {"email": "test@example.com"}, + "server_name": "email_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that no exception is raised + result = await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert result is None + assert mock_pii_guardrail_allow.call_count == 1 + + @pytest.mark.asyncio + async def test_content_guardrail_blocks_pre_call(self, mock_content_guardrail, mock_user_api_key, mock_cache): + """Test that content guardrail properly blocks pre-call""" + proxy_logging = MockProxyLogging([mock_content_guardrail]) + + request_obj = MCPPreCallRequestObject( + tool_name="content_tool", + arguments={"content": "sensitive content"}, + server_name="content_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "content_tool", + "arguments": {"content": "sensitive content"}, + "server_name": "content_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that GuardrailRaisedException is raised + with pytest.raises(GuardrailRaisedException) as excinfo: + await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Verify the error details + assert "Content violates policy" in str(excinfo.value) + assert excinfo.value.guardrail_name == "mock-content-guardrail" + assert mock_content_guardrail.call_count == 1 + + @pytest.mark.asyncio + async def test_http_guardrail_blocks_pre_call(self, mock_http_guardrail, mock_user_api_key, mock_cache): + """Test that HTTP guardrail properly blocks pre-call""" + proxy_logging = MockProxyLogging([mock_http_guardrail]) + + request_obj = MCPPreCallRequestObject( + tool_name="http_tool", + arguments={"url": "http://example.com"}, + server_name="http_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "http_tool", + "arguments": {"url": "http://example.com"}, + "server_name": "http_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that HTTPException is raised + with pytest.raises(HTTPException) as excinfo: + await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Verify the error details + assert excinfo.value.status_code == 400 + assert "Violated guardrail policy" in str(excinfo.value.detail) + assert mock_http_guardrail.call_count == 1 + + @pytest.mark.asyncio + async def test_multiple_guardrails_pre_call(self, mock_pii_guardrail, mock_content_guardrail, mock_user_api_key, mock_cache): + """Test multiple guardrails - first one should block""" + proxy_logging = MockProxyLogging([mock_pii_guardrail, mock_content_guardrail]) + + request_obj = MCPPreCallRequestObject( + tool_name="test_tool", + arguments={"email": "test@example.com"}, + server_name="test_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "test_tool", + "arguments": {"email": "test@example.com"}, + "server_name": "test_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that first guardrail blocks + with pytest.raises(BlockedPiiEntityError): + await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Verify only first guardrail was called + assert mock_pii_guardrail.call_count == 1 + assert mock_content_guardrail.call_count == 0 + + +class TestMCPGuardrailsDuringCall: + """Test MCP guardrails for during-call hooks""" + + @pytest.mark.asyncio + async def test_during_call_guardrail_blocks(self, mock_during_guardrail, mock_user_api_key, mock_cache): + """Test that during-call guardrail properly blocks execution""" + proxy_logging = MockProxyLogging([mock_during_guardrail]) + + request_obj = MCPDuringCallRequestObject( + tool_name="phone_tool", + arguments={"phone": "555-123-4567"}, + server_name="phone_server", + start_time=datetime.now().timestamp(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "phone_tool", + "arguments": {"phone": "555-123-4567"}, + "server_name": "phone_server", + } + + # Test that BlockedPiiEntityError is raised + with pytest.raises(BlockedPiiEntityError) as excinfo: + await proxy_logging.async_during_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Verify the error details + assert excinfo.value.entity_type == "PHONE_NUMBER" + assert excinfo.value.guardrail_name == "mock-during-guardrail" + assert mock_during_guardrail.call_count == 1 + + +class TestMCPGuardrailsIntegration: + """Test MCP guardrails integration with MCP server manager""" + + @pytest.mark.asyncio + async def test_mcp_server_manager_with_guardrails(self): + """Test MCP server manager with guardrail integration""" + + mock_proxy_logging = MockProxyLogging([MockPiiGuardrail(should_block=True)]) + + # Test that guardrail exception is properly raised in the hook + with pytest.raises(BlockedPiiEntityError): + await mock_proxy_logging.async_pre_mcp_tool_call_hook( + kwargs={"name": "email_tool", "arguments": {"email": "test@example.com"}}, + request_obj=MagicMock(), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + @pytest.mark.asyncio + async def test_guardrail_exception_propagation(self): + """Test that guardrail exceptions properly propagate through the system""" + # Test BlockedPiiEntityError + with pytest.raises(BlockedPiiEntityError): + raise BlockedPiiEntityError( + entity_type="EMAIL_ADDRESS", + guardrail_name="test-guardrail" + ) + + # Test GuardrailRaisedException + with pytest.raises(GuardrailRaisedException): + raise GuardrailRaisedException( + guardrail_name="test-guardrail", + message="Test message" + ) + + # Test HTTPException + with pytest.raises(HTTPException): + raise HTTPException( + status_code=400, + detail={"error": "Test error"} + ) + + +class TestMCPGuardrailsErrorHandling: + """Test MCP guardrails error handling scenarios""" + + @pytest.mark.asyncio + async def test_non_guardrail_exception_logging(self, mock_user_api_key, mock_cache): + """Test that non-guardrail exceptions are logged as non-blocking""" + class MockFailingGuardrail(CustomGuardrail): + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + return True + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ): + raise Exception("Non-guardrail error") + + proxy_logging = MockProxyLogging([MockFailingGuardrail()]) + + request_obj = MCPPreCallRequestObject( + tool_name="test_tool", + arguments={"test": "data"}, + server_name="test_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "test_tool", + "arguments": {"test": "data"}, + "server_name": "test_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that non-guardrail exceptions are handled gracefully + result = await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Should return None (not raise exception) + assert result is None + + @pytest.mark.asyncio + async def test_guardrail_should_not_run(self, mock_user_api_key, mock_cache): + """Test that guardrails don't run when should_run_guardrail returns False""" + class MockConditionalGuardrail(CustomGuardrail): + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + return False # Don't run + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ): + raise BlockedPiiEntityError("EMAIL_ADDRESS", "test-guardrail") + + proxy_logging = MockProxyLogging([MockConditionalGuardrail()]) + + request_obj = MCPPreCallRequestObject( + tool_name="test_tool", + arguments={"test": "data"}, + server_name="test_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "test_tool", + "arguments": {"test": "data"}, + "server_name": "test_server", + "user_api_key_auth": mock_user_api_key, + } + + # Test that guardrail doesn't run and no exception is raised + result = await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + # Should return None (guardrail didn't run) + assert result is None + + +class TestMCPGuardrailsEdgeCases: + """Test MCP guardrails edge cases and error conditions""" + + @pytest.mark.asyncio + async def test_empty_guardrails_list(self, mock_user_api_key, mock_cache): + """Test behavior with empty guardrails list""" + proxy_logging = MockProxyLogging([]) # No guardrails + + request_obj = MCPPreCallRequestObject( + tool_name="test_tool", + arguments={"test": "data"}, + server_name="test_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "test_tool", + "arguments": {"test": "data"}, + "server_name": "test_server", + "user_api_key_auth": mock_user_api_key, + } + + # Should return None without any issues + result = await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert result is None + + @pytest.mark.asyncio + async def test_guardrail_with_invalid_data(self, mock_user_api_key, mock_cache): + """Test guardrail behavior with invalid data""" + class MockInvalidDataGuardrail(CustomGuardrail): + def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: + return True + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, + call_type: str, + ): + # Try to access invalid data + invalid_data = data.get("invalid_key", {}) + if invalid_data.get("should_fail"): + raise BlockedPiiEntityError("EMAIL_ADDRESS", "test-guardrail") + return None + + proxy_logging = MockProxyLogging([MockInvalidDataGuardrail()]) + + request_obj = MCPPreCallRequestObject( + tool_name="test_tool", + arguments={"test": "data"}, + server_name="test_server", + user_api_key_auth=mock_user_api_key.model_dump(), + hidden_params=HiddenParams() + ) + + kwargs = { + "name": "test_tool", + "arguments": {"test": "data"}, + "server_name": "test_server", + "user_api_key_auth": mock_user_api_key, + } + + # Should handle invalid data gracefully + result = await proxy_logging.async_pre_mcp_tool_call_hook( + kwargs=kwargs, + request_obj=request_obj, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert result is None + + +if __name__ == "__main__": + pytest.main([__file__]) \ No newline at end of file diff --git a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx index f68decf942f..bce9db2e453 100644 --- a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx @@ -17,7 +17,9 @@ const modeDescriptions = { pre_call: "Before LLM Call - Runs before the LLM call and checks the input (Recommended)", during_call: "During LLM Call - Runs in parallel with the LLM call, with response held until check completes", post_call: "After LLM Call - Runs after the LLM call and checks only the output", - logging_only: "Logging Only - Only runs on logging callbacks without affecting the LLM call" + logging_only: "Logging Only - Only runs on logging callbacks without affecting the LLM call", + pre_mcp_call: "Before MCP Tool Call - Runs before MCP tool execution and validates tool calls", + during_mcp_call: "During MCP Tool Call - Runs in parallel with MCP tool execution for monitoring" }; interface AddGuardrailFormProps { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx b/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx index 6f2845a63db..4727d3713b1 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/ToolTestPanel.tsx @@ -451,7 +451,11 @@ export function ToolTestPanel({ )}
-
{error.message}
+
+                              {(() => {
+                                return error.message;
+                              })()}
+                            
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx index 4500264f54f..78e3c929ca8 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx @@ -161,9 +161,15 @@ const MCPToolsViewer = ({ // Mutation for calling a tool const { mutate: executeTool, isPending: isCallingTool } = useMutation({ - mutationFn: (args: { tool: MCPTool; arguments: Record, authValue: string }) => { + mutationFn: async (args: { tool: MCPTool; arguments: Record, authValue: string }) => { if (!accessToken) throw new Error("Access Token required"); - return callMCPTool(accessToken, args.tool.name, args.arguments, args.authValue); + + try { + const result = await callMCPTool(accessToken, args.tool.name, args.arguments, args.authValue); + return result; + } catch (error) { + throw error; + } }, onSuccess: (data) => { setToolResult(data); diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index c8bc3d2eccb..1443b7f59d6 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5253,6 +5253,7 @@ export const callMCPTool = async ( headers[MCP_AUTH_HEADER] = authValue; } + const response = await fetch(url, { method: "POST", headers, @@ -5263,9 +5264,43 @@ export const callMCPTool = async ( }); if (!response.ok) { - const errorData = await response.text(); - handleError(errorData); - throw new Error("Network response was not ok"); + let errorMessage = "Network response was not ok"; + let errorDetails = null; + + // First, try to get the response as text to see what we're dealing with + const responseText = await response.text(); + + try { + // Try to parse as JSON + const errorData = JSON.parse(responseText); + + if (errorData.detail) { + if (typeof errorData.detail === 'string') { + errorMessage = errorData.detail; + } else if (typeof errorData.detail === 'object') { + errorMessage = errorData.detail.message || errorData.detail.error || "An error occurred"; + errorDetails = errorData.detail; + } + } else { + errorMessage = errorData.message || errorData.error || errorMessage; + } + + } catch (parseError) { + console.error("Failed to parse JSON error response:", parseError); + // If JSON parsing fails, use the raw text + if (responseText) { + errorMessage = responseText; + } + } + + // Create a more informative error object + const enhancedError = new Error(errorMessage); + (enhancedError as any).status = response.status; + (enhancedError as any).statusText = response.statusText; + (enhancedError as any).details = errorDetails; + + handleError(errorMessage); + throw enhancedError; } const data = await response.json(); @@ -5273,6 +5308,11 @@ export const callMCPTool = async ( return data; } catch (error) { console.error("Failed to call MCP tool:", error); + console.error("Error type:", typeof error); + if (error instanceof Error) { + console.error("Error message:", error.message); + console.error("Error stack:", error.stack); + } throw error; } };