From 4469d6d5e37c9b5f17c2045916d3dfba9917d3ca Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 28 May 2026 21:41:37 +0530 Subject: [PATCH] fix(cato): address guardrail review feedback Use proxy-authenticated user identity, forward moderation hook return values, and ensure streaming sender tasks are cancelled and awaited on exit. Co-authored-by: Cursor --- .../cato_networks/cato_networks.py | 101 ++++++--- .../guardrail_hooks/test_cato_networks.py | 198 +++++++++++++++++- 2 files changed, 261 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py index afd61a21223..683df69238e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py @@ -5,6 +5,7 @@ # # +-------------------------------------------------------------+ import asyncio +import contextlib import json import os from typing import TYPE_CHECKING, Any, AsyncGenerator, Optional, Type, Union @@ -12,6 +13,7 @@ from typing import TYPE_CHECKING, Any, AsyncGenerator, Optional, Type, Union from fastapi import HTTPException from pydantic import BaseModel from websockets.asyncio.client import ClientConnection, connect +from websockets.exceptions import ConnectionClosed from litellm import DualCache from litellm._logging import verbose_proxy_logger @@ -64,10 +66,24 @@ class CatoNetworksGuardrail(CustomGuardrail): self.ws_api_base = self.api_base.replace("http://", "ws://").replace( "https://", "wss://" ) - self.dlp_entities: list[dict] = [] - self._max_dlp_entities = 100 super().__init__(**kwargs) + @staticmethod + def _resolve_cato_user_email(user_api_key_dict: UserAPIKeyAuth) -> Optional[str]: + """Use proxy-authenticated identity only; request headers are spoofable.""" + if user_api_key_dict.user_email: + return user_api_key_dict.user_email + if user_api_key_dict.end_user_id: + return str(user_api_key_dict.end_user_id) + return None + + @staticmethod + async def _cancel_background_task(task: asyncio.Task) -> None: + if not task.done(): + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -77,7 +93,10 @@ class CatoNetworksGuardrail(CustomGuardrail): ) -> Union[Exception, str, dict, None]: verbose_proxy_logger.debug("Inside Cato Pre-Call Hook") return await self.call_cato_guardrail( - data, hook="pre_call", key_alias=user_api_key_dict.key_alias + data, + hook="pre_call", + key_alias=user_api_key_dict.key_alias, + user_email=self._resolve_cato_user_email(user_api_key_dict), ) async def async_moderation_hook( @@ -87,18 +106,20 @@ class CatoNetworksGuardrail(CustomGuardrail): call_type: CallTypesLiteral, ) -> Union[Exception, str, dict, None]: verbose_proxy_logger.debug("Inside Cato Moderation Hook") - - await self.call_cato_guardrail( - data, hook="moderation", key_alias=user_api_key_dict.key_alias + return await self.call_cato_guardrail( + data, + hook="moderation", + key_alias=user_api_key_dict.key_alias, + user_email=self._resolve_cato_user_email(user_api_key_dict), ) - return data async def call_cato_guardrail( - self, data: dict, hook: str, key_alias: Optional[str] + self, + data: dict, + hook: str, + key_alias: Optional[str], + user_email: Optional[str] = None, ) -> dict: - user_email = ( - data.get("metadata", {}).get("headers", {}).get("x-cato-user-email") - ) call_id = data.get("litellm_call_id") headers = self._build_cato_headers( hook=hook, @@ -152,11 +173,13 @@ class CatoNetworksGuardrail(CustomGuardrail): return data async def call_cato_guardrail_on_output( - self, request_data: dict, output: str, hook: str, key_alias: Optional[str] + self, + request_data: dict, + output: str, + hook: str, + key_alias: Optional[str], + user_email: Optional[str] = None, ) -> Optional[dict]: - user_email = ( - request_data.get("metadata", {}).get("headers", {}).get("x-cato-user-email") - ) call_id = request_data.get("litellm_call_id") response = await self.async_handler.post( f"{self.api_base}/fw/v1/analyze", @@ -245,7 +268,11 @@ class CatoNetworksGuardrail(CustomGuardrail): ): content = response.choices[0].message.content or "" cato_output_guardrail_result = await self.call_cato_guardrail_on_output( - data, content, hook="output", key_alias=user_api_key_dict.key_alias + data, + content, + hook="output", + key_alias=user_api_key_dict.key_alias, + user_email=self._resolve_cato_user_email(user_api_key_dict), ) if cato_output_guardrail_result and cato_output_guardrail_result.get( "detection_message" @@ -268,9 +295,9 @@ class CatoNetworksGuardrail(CustomGuardrail): response, request_data: dict, ) -> AsyncGenerator[ModelResponseStream, None]: - user_email = ( - request_data.get("metadata", {}).get("headers", {}).get("x-cato-user-email") - ) + from litellm.proxy.proxy_server import StreamingCallbackError + + user_email = self._resolve_cato_user_email(user_api_key_dict) call_id = request_data.get("litellm_call_id") async with connect( f"{self.ws_api_base}/fw/v1/analyze/stream", @@ -284,22 +311,28 @@ class CatoNetworksGuardrail(CustomGuardrail): sender = asyncio.create_task( self.forward_the_stream_to_cato(websocket, response) ) - while True: - result = json.loads(await websocket.recv()) - if verified_chunk := result.get("verified_chunk"): - yield ModelResponseStream.model_validate(verified_chunk) - else: - sender.cancel() - if result.get("done"): + try: + while True: + try: + raw_message = await websocket.recv() + except ConnectionClosed as exc: + raise StreamingCallbackError( + "Cato guardrail connection closed unexpectedly" + ) from exc + result = json.loads(raw_message) + if verified_chunk := result.get("verified_chunk"): + yield ModelResponseStream.model_validate(verified_chunk) + else: + if result.get("done"): + return + if blocking_message := result.get("blocking_message"): + raise StreamingCallbackError(blocking_message) + verbose_proxy_logger.error( + f"Unknown message received from Cato: {result}" + ) return - if blocking_message := result.get("blocking_message"): - from litellm.proxy.proxy_server import StreamingCallbackError - - raise StreamingCallbackError(blocking_message) - verbose_proxy_logger.error( - f"Unknown message received from Cato: {result}" - ) - return + finally: + await self._cancel_background_task(sender) async def forward_the_stream_to_cato( self, diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py index a4c8382faf1..daae0d59a55 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_cato_networks.py @@ -1,10 +1,12 @@ +import json import os import sys -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi.exceptions import HTTPException from httpx import Request, Response +from websockets.exceptions import ConnectionClosed from litellm import DualCache from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import ( @@ -511,11 +513,10 @@ async def test_anonymize_action_without_redacted_chat_returns_data_unchanged(): @pytest.mark.asyncio -async def test_call_cato_guardrail_forwards_user_email_from_metadata(): +async def test_call_cato_guardrail_forwards_user_email_from_auth(): guard = _make_guardrail() data = { "messages": [{"role": "user", "content": "hi"}], - "metadata": {"headers": {"x-cato-user-email": "alice@example.com"}}, "litellm_call_id": "call-xyz", } response = _make_response( @@ -525,7 +526,14 @@ async def test_call_cato_guardrail_forwards_user_email_from_metadata(): "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", return_value=response, ) as mock_post: - await guard.call_cato_guardrail(data, hook="pre_call", key_alias="alias-1") + await guard.async_pre_call_hook( + data=data, + cache=DualCache(), + user_api_key_dict=UserAPIKeyAuth( + key_alias="alias-1", user_email="alice@example.com" + ), + call_type="completion", + ) sent_headers = mock_post.call_args.kwargs["headers"] assert sent_headers["x-cato-user-email"] == "alice@example.com" assert sent_headers["x-cato-call-id"] == "call-xyz" @@ -533,6 +541,47 @@ async def test_call_cato_guardrail_forwards_user_email_from_metadata(): assert sent_headers["x-cato-litellm-hook"] == "pre_call" +@pytest.mark.asyncio +async def test_call_cato_guardrail_ignores_spoofable_metadata_user_email(): + guard = _make_guardrail() + data = { + "messages": [{"role": "user", "content": "hi"}], + "metadata": {"headers": {"x-cato-user-email": "victim@example.com"}}, + } + response = _make_response( + {"analysis_result": {"policy_drill_down": {}}, "required_action": None} + ) + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=response, + ) as mock_post: + await guard.async_pre_call_hook( + data=data, + cache=DualCache(), + user_api_key_dict=UserAPIKeyAuth(user_email="trusted@example.com"), + call_type="completion", + ) + sent_headers = mock_post.call_args.kwargs["headers"] + assert sent_headers["x-cato-user-email"] == "trusted@example.com" + + +@pytest.mark.asyncio +async def test_resolve_cato_user_email_prefers_user_email_over_end_user_id(): + assert ( + CatoNetworksGuardrail._resolve_cato_user_email( + UserAPIKeyAuth(user_email="user@example.com", end_user_id="end-1") + ) + == "user@example.com" + ) + assert ( + CatoNetworksGuardrail._resolve_cato_user_email( + UserAPIKeyAuth(end_user_id="end-1") + ) + == "end-1" + ) + assert CatoNetworksGuardrail._resolve_cato_user_email(UserAPIKeyAuth()) is None + + # ----------------------------------------------------------------------------- # Output-side action branches (call_cato_guardrail_on_output / post_call_success_hook) # ----------------------------------------------------------------------------- @@ -670,3 +719,144 @@ def test_get_config_model_returns_pydantic_class(): ) assert CatoNetworksGuardrail.get_config_model() is CatoNetworksGuardrailConfigModel + + +# ----------------------------------------------------------------------------- +# Streaming hook coverage +# ----------------------------------------------------------------------------- + + +async def _mock_llm_stream(): + yield {"choices": [{"delta": {"content": "hello"}}]} + + +@pytest.mark.asyncio +async def test_streaming_iterator_yields_verified_chunks_and_cancels_sender(): + guard = _make_guardrail() + verified_chunk = { + "id": "chunk-1", + "object": "chat.completion.chunk", + "created": 0, + "model": "gpt-4", + "choices": [{"index": 0, "delta": {"content": "hi"}, "finish_reason": None}], + } + + class MockWebSocket: + recv_calls = 0 + + async def recv(self): + MockWebSocket.recv_calls += 1 + if MockWebSocket.recv_calls == 1: + return json.dumps({"verified_chunk": verified_chunk}) + return json.dumps({"done": True}) + + async def send(self, _chunk): + return None + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return None + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks.connect", + return_value=MockWebSocket(), + ): + chunks = [ + chunk + async for chunk in guard.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(user_email="stream@example.com"), + response=_mock_llm_stream(), + request_data={"litellm_call_id": "stream-call"}, + ) + ] + assert len(chunks) == 1 + assert chunks[0].choices[0].delta.content == "hi" + + +@pytest.mark.asyncio +async def test_streaming_iterator_raises_on_connection_closed(): + guard = _make_guardrail() + from litellm.proxy.proxy_server import StreamingCallbackError + + class ClosedWebSocket: + async def recv(self): + raise ConnectionClosed(None, None) + + async def send(self, _chunk): + return None + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return None + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks.connect", + return_value=ClosedWebSocket(), + ): + with pytest.raises( + StreamingCallbackError, match="connection closed unexpectedly" + ): + async for _ in guard.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=_mock_llm_stream(), + request_data={}, + ): + pass + + +@pytest.mark.asyncio +async def test_streaming_iterator_raises_on_blocking_message(): + guard = _make_guardrail() + from litellm.proxy.proxy_server import StreamingCallbackError + + class BlockingWebSocket: + async def recv(self): + return json.dumps({"blocking_message": "blocked by policy"}) + + async def send(self, _chunk): + return None + + async def __aenter__(self): + return self + + async def __aexit__(self, exc_type, exc, tb): + return None + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks.connect", + return_value=BlockingWebSocket(), + ): + with pytest.raises(StreamingCallbackError, match="blocked by policy"): + async for _ in guard.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=_mock_llm_stream(), + request_data={}, + ): + pass + + +@pytest.mark.asyncio +async def test_forward_the_stream_to_cato_serializes_chunks(): + guard = _make_guardrail() + websocket = MagicMock() + websocket.send = AsyncMock() + + async def response_iter(): + yield {"role": "assistant"} + yield ModelResponse( + choices=[ + { + "finish_reason": "stop", + "index": 0, + "message": {"content": "done", "role": "assistant"}, + } + ] + ) + + await guard.forward_the_stream_to_cato(websocket, response_iter()) + assert websocket.send.await_count == 3 + assert json.loads(websocket.send.await_args_list[-1].args[0]) == {"done": True}