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 <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-28 21:41:37 +05:30
parent dbba9f1bb3
commit 4469d6d5e3
No known key found for this signature in database
2 changed files with 261 additions and 38 deletions

View file

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

View file

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