diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 2e91e082bd4..8facd913ddf 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c8ff6e262d2..8274747bc28 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 49a1f1314f0..7990d7922a9 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 5f22ca021ac..d359238bf45 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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, ), diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 8448db11904..bfcf808d218 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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: diff --git a/tests/test_litellm/integrations/test_guardrail_logging_sync.py b/tests/test_litellm/integrations/test_guardrail_logging_sync.py index f9e1a3efbd0..2547afafafa 100644 --- a/tests/test_litellm/integrations/test_guardrail_logging_sync.py +++ b/tests/test_litellm/integrations/test_guardrail_logging_sync.py @@ -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() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index b56a12db5b1..509f92c102b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -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"} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 850d01c6e34..f8af3e0c417 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 54fd5242d5f..3fc53ea9763 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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 diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index a8e81e92ebd..d98f6d284e7 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -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 diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py index 9defb309863..462f7c65327 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_mcp_bridging.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index d60fff66c44..f1f48c90f4b 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -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