mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
[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:
parent
c125ae453b
commit
900c7f45c0
38 changed files with 1294 additions and 104 deletions
|
|
@ -173,6 +173,7 @@ class AporiaGuardrail(CustomGuardrail):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ class _ENTERPRISE_GoogleTextModeration(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
text = ""
|
||||
|
|
|
|||
|
|
@ -105,6 +105,7 @@ class _ENTERPRISE_LlamaGuard(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -127,6 +127,7 @@ class _ENTERPRISE_LLMGuard(CustomLogger):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -147,6 +147,7 @@ class PagerDutyAlerting(SlackAlerting):
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -188,6 +188,7 @@ class AporiaGuardrail(CustomGuardrail):
|
|||
"moderation",
|
||||
"audio_transcription",
|
||||
"responses",
|
||||
"mcp_call",
|
||||
],
|
||||
):
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
|
|
|
|||
|
|
@ -123,6 +123,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -218,6 +218,7 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -204,6 +204,7 @@ class GuardrailsAI(CustomGuardrail):
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[
|
||||
Union[Exception, str, dict]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -245,6 +245,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -197,6 +197,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
"audio_transcription",
|
||||
"pass_through_endpoint",
|
||||
"rerank",
|
||||
"mcp_call",
|
||||
],
|
||||
) -> Optional[
|
||||
Union[Exception, str, dict]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
734
tests/mcp_tests/test_mcp_guardrails.py
Normal file
734
tests/mcp_tests/test_mcp_guardrails.py
Normal 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__])
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue