[MCP Gateway] Litellm mcp pre and during guardrails (#13188)

* add guardrail support

* add guardrail support

* guardrails for MCP

* added changes

* add mcp guardrails

* added test

* add ui

* fix guardrail form

* working with cursor

* remvoe print

* fix mcp servertests

* fix mypy and remove console logs

* fix mypy and remove console logs

* fix mypy tests
This commit is contained in:
Jugal D. Bhatt 2025-08-01 20:02:25 -07:00 • committed by GitHub
parent c125ae453b
commit 900c7f45c0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
38 changed files with 1294 additions and 104 deletions

View file

@ -173,6 +173,7 @@ class AporiaGuardrail(CustomGuardrail):
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
):
from litellm.proxy.common_utils.callback_utils import (

View file

@ -95,6 +95,7 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger):
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
):
"""

View file

@ -42,6 +42,7 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
):
text = ""

View file

@ -105,6 +105,7 @@ class _ENTERPRISE_LlamaGuard(CustomLogger):
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
):
"""

View file

@ -127,6 +127,7 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
):
"""

View file

@ -147,6 +147,7 @@ class PagerDutyAlerting(SlackAlerting):
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[Union[Exception, str, dict]]:
"""

View file

@ -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]:
"""

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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",
],
):
"""

View file

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

View file

@ -188,6 +188,7 @@ class AporiaGuardrail(CustomGuardrail):
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
):
from litellm.proxy.common_utils.callback_utils import (

View file

@ -123,6 +123,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[Dict[str, Any]]:
"""

View file

@ -218,6 +218,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[Dict[str, Any]]:
"""

View file

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

View file

@ -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",
],
):
"""

View file

@ -204,6 +204,7 @@ class GuardrailsAI(CustomGuardrail):
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[
Union[Exception, str, dict]

View file

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

View file

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

View file

@ -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",
],
):
"""

View file

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

View file

@ -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]]:
"""

View file

@ -245,6 +245,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[Dict[str, Any]]:
"""

View file

@ -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]]:
"""

View file

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

View file

@ -197,6 +197,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[
Union[Exception, str, dict]

View file

@ -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: <tool_name>\nArguments: <json_arguments>"
# 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)

View file

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

View file

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

View file

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

View file

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

View file

@ -451,7 +451,11 @@ export function ToolTestPanel({
)}
</div>
<div className="bg-white border border-red-200 rounded p-2 max-h-48 overflow-y-auto">
<pre className="text-xs whitespace-pre-wrap text-red-700 font-mono">{error.message}</pre>
<pre className="text-xs whitespace-pre-wrap text-red-700 font-mono">
{(() => {
return error.message;
})()}
</pre>
</div>
</div>
</div>

View file

@ -161,9 +161,15 @@ const MCPToolsViewer = ({
// Mutation for calling a tool
const { mutate: executeTool, isPending: isCallingTool } = useMutation({
mutationFn: (args: { tool: MCPTool; arguments: Record<string, any>, authValue: string }) => {
mutationFn: async (args: { tool: MCPTool; arguments: Record<string, any>, 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);

View file

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