mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(guardrails): persist guardrail_information for MCP tool calls
This commit is contained in:
parent
f6b9518ddb
commit
f5b92ab6f1
12 changed files with 538 additions and 9 deletions
|
|
@ -1241,16 +1241,22 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object)
|
|||
that does not share identity with the one in request_data. This helper
|
||||
bridges that gap so guardrail_information is non-null in spend logs for all
|
||||
routes, not just /v1/chat/completions.
|
||||
|
||||
Hooks whose signature carries no ``logging_obj`` (``async_pre_call_hook``,
|
||||
``async_moderation_hook``) leave it ``None``; for those the object is taken from
|
||||
``request_data["litellm_logging_obj"]``, which every route that builds its own
|
||||
request dict already seeds.
|
||||
"""
|
||||
if logging_obj is None:
|
||||
target: Final = logging_obj if logging_obj is not None else request_data.get("litellm_logging_obj")
|
||||
if target is None:
|
||||
return
|
||||
meta_src: Final = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
|
||||
slg_info: Final = meta_src.get("standard_logging_guardrail_information")
|
||||
if not slg_info:
|
||||
return
|
||||
entries: Final[list] = slg_info if isinstance(slg_info, list) else [slg_info]
|
||||
mcd: Final = getattr(logging_obj, "model_call_details", None) or {}
|
||||
_append_slg_to_litellm_params(getattr(logging_obj, "litellm_params", None), entries)
|
||||
mcd: Final = getattr(target, "model_call_details", None) or {}
|
||||
_append_slg_to_litellm_params(getattr(target, "litellm_params", None), entries)
|
||||
_append_slg_to_litellm_params(mcd.get("litellm_params"), entries)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -161,6 +161,7 @@ if TYPE_CHECKING:
|
|||
from mcp.types import CreateMessageRequestParams
|
||||
|
||||
from litellm.caching.caching import InMemoryCache
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.mcp_server.mcp_toolset import MCPToolset
|
||||
|
||||
try:
|
||||
|
|
@ -4542,6 +4543,7 @@ class MCPServerManager:
|
|||
proxy_logging_obj: ProxyLogging | None,
|
||||
server: MCPServer,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Run pre-call checks and guardrail hooks for an MCP tool call.
|
||||
|
|
@ -4609,7 +4611,9 @@ class MCPServerManager:
|
|||
mcp_request_obj: Final = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
|
||||
|
||||
# Convert to LLM format for existing guardrail compatibility
|
||||
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(mcp_request_obj, pre_hook_kwargs)
|
||||
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(
|
||||
mcp_request_obj, pre_hook_kwargs, litellm_logging_obj=litellm_logging_obj
|
||||
)
|
||||
|
||||
try:
|
||||
# Use standard pre_call_hook
|
||||
|
|
@ -4645,6 +4649,7 @@ class MCPServerManager:
|
|||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
start_time: datetime.datetime,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
):
|
||||
"""Create and return a during hook task for MCP tool calls."""
|
||||
from litellm.types.llms.base import HiddenParams
|
||||
|
|
@ -4665,7 +4670,9 @@ class MCPServerManager:
|
|||
"user_api_key_auth": user_api_key_auth,
|
||||
}
|
||||
|
||||
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, during_hook_kwargs)
|
||||
synthetic_llm_data: Final = proxy_logging_obj._convert_mcp_to_llm_format(
|
||||
request_obj, during_hook_kwargs, litellm_logging_obj=litellm_logging_obj
|
||||
)
|
||||
|
||||
return asyncio.create_task(
|
||||
proxy_logging_obj.during_call_hook(
|
||||
|
|
@ -5202,6 +5209,7 @@ class MCPServerManager:
|
|||
oauth2_headers: dict[str, str] | None = None,
|
||||
raw_headers: dict[str, str] | None = None,
|
||||
host_progress_callback: Callable | None = None,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
|
|
@ -5244,6 +5252,7 @@ class MCPServerManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"]
|
||||
|
|
@ -5258,6 +5267,7 @@ class MCPServerManager:
|
|||
user_api_key_auth=user_api_key_auth,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
start_time=start_time,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
tasks.append(during_hook_task)
|
||||
|
||||
|
|
|
|||
|
|
@ -2824,6 +2824,7 @@ if MCP_AVAILABLE:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
server=mcp_server,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
# `pre_call_tool_check` may return guardrail-modified
|
||||
# arguments; honor them on the local path too.
|
||||
|
|
@ -2962,6 +2963,7 @@ if MCP_AVAILABLE:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
server=prefix_server,
|
||||
raw_headers=raw_headers,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
if "arguments" in hook_result:
|
||||
arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args
|
||||
|
|
@ -3000,6 +3002,7 @@ if MCP_AVAILABLE:
|
|||
litellm_logging_obj.model_call_details if litellm_logging_obj is not None else dict(request_data)
|
||||
),
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
_MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset(
|
||||
|
|
@ -3326,6 +3329,7 @@ if MCP_AVAILABLE:
|
|||
raw_headers=raw_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
host_progress_callback=host_progress_callback,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
|
||||
return call_tool_result
|
||||
|
|
|
|||
|
|
@ -639,9 +639,18 @@ class ProxyLogging:
|
|||
return user_api_key_auth_obj.__dict__
|
||||
return {}
|
||||
|
||||
def _convert_mcp_to_llm_format(self, request_obj, kwargs: dict) -> dict:
|
||||
def _convert_mcp_to_llm_format(
|
||||
self,
|
||||
request_obj,
|
||||
kwargs: dict,
|
||||
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Convert MCP tool call to LLM message format for existing guardrail validation.
|
||||
|
||||
``litellm_logging_obj`` is the request's logging object, seeded onto the synthetic
|
||||
request so a guardrail's recorded evaluation reaches the object that builds the
|
||||
spend-log payload — the MCP equivalent of what every LLM route already carries.
|
||||
"""
|
||||
from litellm.types.llms.openai import ChatCompletionUserMessage
|
||||
|
||||
|
|
@ -670,6 +679,7 @@ class ProxyLogging:
|
|||
# (e.g. MCPJWTSigner) to independently verify the caller's identity
|
||||
# before re-signing an outbound token (FR-5 verify+re-sign).
|
||||
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
|
||||
"litellm_logging_obj": litellm_logging_obj,
|
||||
}
|
||||
|
||||
return synthetic_data
|
||||
|
|
@ -2487,10 +2497,14 @@ class ProxyLogging:
|
|||
response: "CallToolResult",
|
||||
request_data: Mapping[str, Any],
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> "CallToolResult":
|
||||
"""
|
||||
Run guardrails configured for ``post_mcp_call`` against an MCP tool result.
|
||||
|
||||
``litellm_logging_obj`` is supplied by the caller rather than read out of
|
||||
``request_data``, which on this path is ``model_call_details`` and never carries it.
|
||||
|
||||
The MCP counterpart of ``post_call_success_hook``: guardrails that
|
||||
implement ``apply_guardrail`` see the tool result's text through the
|
||||
unified guardrail seam (``MCPGuardrailTranslationHandler``), so a text
|
||||
|
|
@ -2527,7 +2541,7 @@ class ProxyLogging:
|
|||
handler_cls().process_output_response(
|
||||
response=response,
|
||||
guardrail_to_apply=callback,
|
||||
litellm_logging_obj=request_data.get("litellm_logging_obj"),
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -789,6 +789,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
if proxy_logging_obj:
|
||||
|
|
@ -800,6 +801,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
else {"mcp_tool_name": tool_name}
|
||||
),
|
||||
user_api_key_dict=user_api_key_auth,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
if litellm_logging_obj:
|
||||
|
|
|
|||
|
|
@ -122,6 +122,49 @@ def test_noop_when_logging_obj_is_none():
|
|||
_sync_guardrail_info_to_logging_obj(request_data, None)
|
||||
|
||||
|
||||
def test_syncs_using_logging_obj_from_request_data():
|
||||
"""Hooks with no logging_obj parameter (async_pre_call_hook) leave the decorator's
|
||||
kwarg None; the object must then come off request_data or direct-hook guardrails
|
||||
record nothing on the MCP path."""
|
||||
entry = _make_slg_entry()
|
||||
logging_obj = _FakeLogging()
|
||||
request_data = {
|
||||
"metadata": {"standard_logging_guardrail_information": [entry]},
|
||||
"litellm_logging_obj": logging_obj,
|
||||
}
|
||||
|
||||
_sync_guardrail_info_to_logging_obj(request_data, None)
|
||||
|
||||
assert logging_obj.litellm_params["metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
] == [entry]
|
||||
|
||||
|
||||
def test_explicit_logging_obj_wins_over_request_data():
|
||||
entry = _make_slg_entry()
|
||||
explicit, on_data = _FakeLogging(), _FakeLogging()
|
||||
request_data = {
|
||||
"metadata": {"standard_logging_guardrail_information": [entry]},
|
||||
"litellm_logging_obj": on_data,
|
||||
}
|
||||
|
||||
_sync_guardrail_info_to_logging_obj(request_data, explicit)
|
||||
|
||||
assert explicit.litellm_params["metadata"][
|
||||
"standard_logging_guardrail_information"
|
||||
] == [entry]
|
||||
assert (
|
||||
"standard_logging_guardrail_information" not in on_data.litellm_params["metadata"]
|
||||
)
|
||||
|
||||
|
||||
def test_no_logging_obj_anywhere_is_a_noop():
|
||||
_sync_guardrail_info_to_logging_obj(
|
||||
{"metadata": {"standard_logging_guardrail_information": [_make_slg_entry()]}},
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def test_writes_to_model_call_details_too():
|
||||
"""Also writes into model_call_details["litellm_params"]["metadata"]."""
|
||||
entry = _make_slg_entry()
|
||||
|
|
|
|||
|
|
@ -783,7 +783,7 @@ class TestMcpRateLimitServerNameSurfacing:
|
|||
|
||||
captured = {}
|
||||
|
||||
def capture_convert(request_obj, kwargs):
|
||||
def capture_convert(request_obj, kwargs, litellm_logging_obj=None):
|
||||
captured["kwargs"] = kwargs
|
||||
return {"model": "fake"}
|
||||
|
||||
|
|
|
|||
|
|
@ -7977,6 +7977,7 @@ async def test_post_mcp_call_guardrails_return_the_rewritten_result():
|
|||
hook_kwargs = proxy_logging_mock.post_mcp_call_hook.await_args.kwargs
|
||||
assert hook_kwargs["response"] is raw_result
|
||||
assert hook_kwargs["request_data"] is logging_obj.model_call_details
|
||||
assert hook_kwargs["litellm_logging_obj"] is logging_obj
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -8236,3 +8237,157 @@ class TestListFiltersHonorThePrefixBoundary:
|
|||
|
||||
assert listed == callable_, f"grants={grants!r} listed={listed} callable={callable_}"
|
||||
assert listed is expected, f"grants={grants!r} expected={expected} got={listed}"
|
||||
|
||||
|
||||
class TestExecuteMCPToolForwardsLoggingObject:
|
||||
"""Every dispatch branch must hand the request's logging object downstream.
|
||||
|
||||
The guardrail hooks record their evaluation into the object they are given; a branch
|
||||
that drops it writes the record into a throwaway and the spend-log row for that tool
|
||||
call reports guardrail_information as null.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _fake_server(name="guardrail-info-server"):
|
||||
fake_server = MagicMock()
|
||||
fake_server.name = name
|
||||
fake_server.is_byok = False
|
||||
fake_server.auth_type = None
|
||||
fake_server.mcp_info = None
|
||||
fake_server.server_id = "srv-guardrail-info"
|
||||
fake_server.server_name = name
|
||||
fake_server.alias = None
|
||||
fake_server.short_prefix = None
|
||||
return fake_server
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_prefixed_tool_branch_forwards_logging_obj(self):
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
fake_server = self._fake_server()
|
||||
sentinel = _mock_mcp_logging_obj()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_get_mcp_server_from_tool_name",
|
||||
return_value=fake_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"pre_call_tool_check",
|
||||
new=AsyncMock(return_value={}),
|
||||
) as pre_check,
|
||||
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool",
|
||||
new=AsyncMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._run_post_mcp_call_guardrails",
|
||||
new=AsyncMock(side_effect=lambda result, **_: result),
|
||||
),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="list_pets",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[fake_server],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="u-1"),
|
||||
litellm_logging_obj=sentinel,
|
||||
)
|
||||
|
||||
assert pre_check.await_args.kwargs["litellm_logging_obj"] is sentinel
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unprefixed_local_registry_fallback_forwards_logging_obj(self):
|
||||
"""The prefix-stripped local-registry fallback: no server resolves from the tool
|
||||
name, so the branch re-resolves the server from allowed_mcp_servers and runs the
|
||||
pre-call check itself. It must forward the logging object like the others."""
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
fake_server = self._fake_server(name="legacysrv")
|
||||
sentinel = _mock_mcp_logging_obj()
|
||||
|
||||
def _get_tool(tool_name):
|
||||
return MagicMock() if tool_name == "echo" else None
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_get_mcp_server_from_tool_name",
|
||||
return_value=None,
|
||||
),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"pre_call_tool_check",
|
||||
new=AsyncMock(return_value={}),
|
||||
) as pre_check,
|
||||
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", side_effect=_get_tool),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool",
|
||||
new=AsyncMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._run_post_mcp_call_guardrails",
|
||||
new=AsyncMock(side_effect=lambda result, **_: result),
|
||||
),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="legacysrv-echo",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[fake_server],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="u-1"),
|
||||
litellm_logging_obj=sentinel,
|
||||
)
|
||||
|
||||
assert pre_check.await_args.kwargs["litellm_logging_obj"] is sentinel
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_server_branch_forwards_logging_obj(self):
|
||||
"""The customer's route: a managed HTTP MCP server dispatched through call_tool."""
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_module
|
||||
|
||||
fake_server = self._fake_server()
|
||||
sentinel = _mock_mcp_logging_obj()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"_get_mcp_server_from_tool_name",
|
||||
return_value=fake_server,
|
||||
),
|
||||
patch.object(
|
||||
mcp_module.global_mcp_server_manager,
|
||||
"call_tool",
|
||||
new=AsyncMock(return_value=_call_tool_result(False, "ok")),
|
||||
) as call_tool,
|
||||
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed",
|
||||
return_value=True,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._run_post_mcp_call_guardrails",
|
||||
new=AsyncMock(side_effect=lambda result, **_: result),
|
||||
),
|
||||
):
|
||||
await mcp_module.execute_mcp_tool(
|
||||
name="list_pets",
|
||||
arguments={},
|
||||
allowed_mcp_servers=[fake_server],
|
||||
start_time=datetime.now(),
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="u-1"),
|
||||
litellm_logging_obj=sentinel,
|
||||
)
|
||||
|
||||
assert call_tool.await_args.kwargs["litellm_logging_obj"] is sentinel
|
||||
|
|
|
|||
|
|
@ -9793,3 +9793,196 @@ class TestToolAuthorizationIsNotConditionalOnLogging:
|
|||
)
|
||||
|
||||
upstream.assert_awaited_once()
|
||||
|
||||
|
||||
class TestMCPGuardrailInformationReachesLoggingObject:
|
||||
"""MCP guardrail evaluations must land on the request's logging object.
|
||||
|
||||
The spend-log payload is built from ``logging_obj.litellm_params["metadata"]``.
|
||||
The guardrail writes its record into the synthetic request dict the MCP gateway
|
||||
hands it, so unless that dict carries ``litellm_logging_obj`` the record is written
|
||||
into a throwaway object and MCP rows report ``guardrail_information: null``.
|
||||
"""
|
||||
|
||||
@pytest.fixture
|
||||
def restore_callbacks(self):
|
||||
import litellm
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
original = list(litellm.callbacks)
|
||||
yield
|
||||
litellm.callbacks = original
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
||||
@staticmethod
|
||||
def _make_logging_obj():
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
logging_obj = Logging(
|
||||
model="MCP: search",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="call_mcp_tool",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="mcp-guardrail-info-test",
|
||||
function_id="mcp-guardrail-info-test",
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
litellm_params={"metadata": {}},
|
||||
optional_params={},
|
||||
)
|
||||
assert logging_obj.model_call_details["litellm_params"] is logging_obj.litellm_params
|
||||
return logging_obj
|
||||
|
||||
@staticmethod
|
||||
def _recorded_entries(logging_obj):
|
||||
return logging_obj.litellm_params["metadata"].get(
|
||||
"standard_logging_guardrail_information", []
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _make_direct_hook_guardrail(event_hook, hook_name):
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
|
||||
class _DirectHookGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="mcp-direct-hook-guardrail",
|
||||
event_hook=event_hook,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
return data
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_moderation_hook(self, data, user_api_key_dict, call_type):
|
||||
return data
|
||||
|
||||
guardrail = _DirectHookGuardrail()
|
||||
assert hook_name in type(guardrail).__dict__
|
||||
return guardrail
|
||||
|
||||
@staticmethod
|
||||
def _make_open_server():
|
||||
return MCPServer(
|
||||
server_id="guardrail-info-server",
|
||||
name="guardrail-info-server",
|
||||
transport=MCPTransport.stdio,
|
||||
allowed_tools=None,
|
||||
disallowed_tools=None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _make_user_api_key_auth():
|
||||
return UserAPIKeyAuth(
|
||||
api_key="sk-guardrail-info-test",
|
||||
user_id="u-1",
|
||||
team_id="t-1",
|
||||
object_permission=None,
|
||||
object_permission_id=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_tool_check_records_guardrail_info_on_logging_obj(self, restore_callbacks):
|
||||
"""A pre_mcp_call guardrail's evaluation must reach the logging object that builds the
|
||||
spend-log payload. Drop litellm_logging_obj at any hop between pre_call_tool_check and
|
||||
the guardrail and this fails with guardrail_information null, which is the bug."""
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
guardrail = self._make_direct_hook_guardrail(
|
||||
GuardrailEventHooks.pre_mcp_call, "async_pre_call_hook"
|
||||
)
|
||||
litellm.callbacks = [guardrail]
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
logging_obj = self._make_logging_obj()
|
||||
|
||||
await MCPServerManager().pre_call_tool_check(
|
||||
name="search",
|
||||
arguments={"q": "hello"},
|
||||
server_name="guardrail-info-server",
|
||||
user_api_key_auth=self._make_user_api_key_auth(),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
server=self._make_open_server(),
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
entries = self._recorded_entries(logging_obj)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == guardrail.guardrail_name
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_during_hook_records_guardrail_info_on_logging_obj(self, restore_callbacks):
|
||||
"""during_mcp_call runs on the same synthetic dict and must record the same way."""
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
guardrail = self._make_direct_hook_guardrail(
|
||||
GuardrailEventHooks.during_mcp_call, "async_moderation_hook"
|
||||
)
|
||||
litellm.callbacks = [guardrail]
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
logging_obj = self._make_logging_obj()
|
||||
|
||||
await MCPServerManager()._create_during_hook_task(
|
||||
name="search",
|
||||
arguments={"q": "hello"},
|
||||
server_name_from_prefix="guardrail-info-server",
|
||||
user_api_key_auth=self._make_user_api_key_auth(),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
start_time=datetime.now(),
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
entries = self._recorded_entries(logging_obj)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == guardrail.guardrail_name
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_tool_forwards_logging_obj_to_both_hooks(self):
|
||||
"""call_tool is the single funnel for the managed-server route; it must hand the
|
||||
logging object to both the pre-call check and the during-call task."""
|
||||
manager = MCPServerManager()
|
||||
manager.tool_name_to_mcp_server_name_mapping = {}
|
||||
sentinel = object()
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
manager,
|
||||
"_resolve_mcp_server_for_tool_call",
|
||||
return_value=self._make_open_server(),
|
||||
),
|
||||
patch.object(
|
||||
manager, "pre_call_tool_check", new=AsyncMock(return_value={})
|
||||
) as pre_check,
|
||||
patch.object(manager, "_create_during_hook_task") as during_task,
|
||||
patch.object(
|
||||
manager,
|
||||
"_call_regular_mcp_tool",
|
||||
new=AsyncMock(return_value=CallToolResult(content=[], isError=False)),
|
||||
),
|
||||
):
|
||||
during_task.return_value = MagicMock()
|
||||
await manager.call_tool(
|
||||
server_name="guardrail-info-server",
|
||||
name="search",
|
||||
arguments={"q": "hello"},
|
||||
user_api_key_auth=self._make_user_api_key_auth(),
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=sentinel,
|
||||
)
|
||||
|
||||
assert pre_check.await_args.kwargs["litellm_logging_obj"] is sentinel
|
||||
assert during_task.call_args.kwargs["litellm_logging_obj"] is sentinel
|
||||
|
|
|
|||
|
|
@ -7,7 +7,10 @@ import pytest
|
|||
from fastapi import HTTPException
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.proxy._types import ProxyErrorTypes
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
|
@ -1169,3 +1172,52 @@ async def test_prisma_health_check_failure_redacts_database_credentials(caplog):
|
|||
assert emitted
|
||||
assert all("hunter2" not in message for message in emitted)
|
||||
assert any("postgresql://REDACTED@db.internal" in message for message in emitted)
|
||||
|
||||
|
||||
class _DecoratedMCPGuardrail(_RecordingMCPGuardrail):
|
||||
"""Unified guardrail whose apply_guardrail auto-records, as presidio and noma_v2 do."""
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
||||
return await _RecordingMCPGuardrail.apply_guardrail(
|
||||
self, inputs=inputs, request_data=request_data, input_type=input_type, **kwargs
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_mcp_call_hook_records_guardrail_info_on_logging_obj(restore_callbacks):
|
||||
"""post_mcp_call used to read the logging object out of request_data, a key
|
||||
model_call_details never carries, so the result scan was never recorded."""
|
||||
from datetime import datetime
|
||||
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
guardrail = _DecoratedMCPGuardrail(event_hook=GuardrailEventHooks.post_mcp_call)
|
||||
litellm.callbacks = [guardrail]
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
|
||||
logging_obj = Logging(
|
||||
model="MCP: echo",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="call_mcp_tool",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="post-mcp-guardrail-info-test",
|
||||
function_id="post-mcp-guardrail-info-test",
|
||||
)
|
||||
logging_obj.update_environment_variables(litellm_params={"metadata": {}}, optional_params={})
|
||||
result = CallToolResult(content=[TextContent(type="text", text="jane@example.com")], isError=False)
|
||||
|
||||
await proxy_logging_obj.post_mcp_call_hook(
|
||||
response=result,
|
||||
request_data=logging_obj.model_call_details,
|
||||
user_api_key_dict=None,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
entries = logging_obj.litellm_params["metadata"].get("standard_logging_guardrail_information", [])
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == guardrail.guardrail_name
|
||||
|
|
|
|||
|
|
@ -81,6 +81,26 @@ def test_convert_mcp_to_llm_format_missing_request_obj_raises(proxy_logging):
|
|||
proxy_logging._convert_mcp_to_llm_format(request_obj=None, kwargs={})
|
||||
|
||||
|
||||
def test_convert_mcp_to_llm_format_seeds_logging_obj(proxy_logging, make_mcp_request_obj):
|
||||
"""The synthetic guardrail request must carry the request's logging object, otherwise
|
||||
the guardrail's recorded evaluation has nowhere to land and MCP spend-log rows report
|
||||
guardrail_information as null."""
|
||||
sentinel = object()
|
||||
out = proxy_logging._convert_mcp_to_llm_format(
|
||||
request_obj=make_mcp_request_obj(tool_name="search", arguments={}),
|
||||
kwargs={},
|
||||
litellm_logging_obj=sentinel,
|
||||
)
|
||||
assert out["litellm_logging_obj"] is sentinel
|
||||
|
||||
|
||||
def test_convert_mcp_to_llm_format_logging_obj_defaults_to_none(proxy_logging, make_mcp_request_obj):
|
||||
out = proxy_logging._convert_mcp_to_llm_format(
|
||||
request_obj=make_mcp_request_obj(tool_name="search", arguments={}), kwargs={}
|
||||
)
|
||||
assert out["litellm_logging_obj"] is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _convert_llm_result_to_mcp_response
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -648,3 +648,33 @@ def test_extract_tool_call_details_still_prefers_openai_arguments():
|
|||
assert name == "get_weather"
|
||||
assert call_id == "call_123"
|
||||
assert arguments == '{"city": "Paris"}'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_tool_calls_threads_logging_obj_into_call_tool(monkeypatch):
|
||||
"""The Responses-API MCP path builds its own logging object; call_tool must receive it
|
||||
or the pre/during guardrail records for these tool calls are written into a throwaway
|
||||
and the spend-log row reports guardrail_information as null."""
|
||||
call_tool_mock = _setup_mcp_call_environment(monkeypatch)
|
||||
sentinel_logging_obj = MagicMock()
|
||||
sentinel_logging_obj.model_call_details = {}
|
||||
monkeypatch.setattr(
|
||||
sys.modules["litellm.responses.mcp.litellm_proxy_mcp_handler"],
|
||||
"function_setup",
|
||||
MagicMock(return_value=(sentinel_logging_obj, {})),
|
||||
)
|
||||
tool_name = "read_wiki_structure"
|
||||
tool_calls = [
|
||||
{
|
||||
"id": "call-1",
|
||||
"function": {"name": tool_name, "arguments": "{}"},
|
||||
}
|
||||
]
|
||||
|
||||
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map={tool_name: "deepwiki"},
|
||||
tool_calls=tool_calls,
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
||||
assert call_tool_mock.await_args.kwargs["litellm_logging_obj"] is sentinel_logging_obj
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue