mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
dbba9f1bb3
commit
4469d6d5e3
2 changed files with 261 additions and 38 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue