mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails/peyeeye): address review — shared cache, stricter errors
- Use module-level `dc` (`global_cache`) from `custom_guardrail` for both pre/post hooks, matching `LassoGuardrail`. Pre-call wrote to the param `cache` while post-call read `litellm.cache`, which silently skipped rehydration whenever `litellm.cache` was None and leaked placeholders (e.g. `[EMAIL_1]`) to callers. - Drop the unused `fastapi.HTTPException` branch in `_reraise_api_error` (httpx errors are never `HTTPException`); annotate the helper as `NoReturn` so static analysers can see `_post`'s control flow. - Raise `PeyeeyeGuardrailAPIError` instead of silently passing the unredacted source through when `/v1/redact` returns an unexpected shape — never forward sensitive content on a parse fallback. - Add an end-to-end test that exercises pre-call → post-call through the shared cache; tighten the existing pre-call tests to assert on `global_cache` rather than the locally-passed `DualCache`.
This commit is contained in:
parent
a636014380
commit
3eda6087d1
2 changed files with 39 additions and 25 deletions
|
|
@ -12,6 +12,7 @@ from typing import (
|
|||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
NoReturn,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
|
|
@ -25,13 +26,12 @@ except ImportError:
|
|||
httpx = None # type: ignore
|
||||
HTTPX_AVAILABLE = False
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
dc as global_cache,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
@ -153,7 +153,7 @@ class PeyeeyeGuardrail(CustomGuardrail):
|
|||
if session_id:
|
||||
cache_key = self._cache_key(data)
|
||||
try:
|
||||
cache.set_cache(
|
||||
global_cache.set_cache(
|
||||
cache_key, session_id, ttl=SESSION_CACHE_TTL_SECONDS
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
@ -177,7 +177,7 @@ class PeyeeyeGuardrail(CustomGuardrail):
|
|||
|
||||
cache_key = self._cache_key(data)
|
||||
try:
|
||||
session_id = litellm.cache.get_cache(cache_key) if litellm.cache else None
|
||||
session_id = global_cache.get_cache(cache_key)
|
||||
except Exception:
|
||||
session_id = None
|
||||
if not session_id:
|
||||
|
|
@ -214,8 +214,7 @@ class PeyeeyeGuardrail(CustomGuardrail):
|
|||
"peyeeye: best-effort session cleanup failed: %s", e
|
||||
)
|
||||
try:
|
||||
if litellm.cache:
|
||||
litellm.cache.delete_cache(cache_key)
|
||||
global_cache.delete_cache(cache_key)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
|
@ -249,7 +248,10 @@ class PeyeeyeGuardrail(CustomGuardrail):
|
|||
elif isinstance(out_text, list):
|
||||
redacted = [str(x) for x in out_text]
|
||||
else:
|
||||
redacted = list(texts) # fallback
|
||||
raise PeyeeyeGuardrailAPIError(
|
||||
"peyeeye /v1/redact returned unexpected response shape; "
|
||||
"refusing to forward unredacted text"
|
||||
)
|
||||
|
||||
if self.peyeeye_session_mode == "stateless":
|
||||
session_id = payload.get("rehydration_key")
|
||||
|
|
@ -291,9 +293,7 @@ class PeyeeyeGuardrail(CustomGuardrail):
|
|||
}
|
||||
|
||||
@staticmethod
|
||||
def _reraise_api_error(error: Exception, path: str) -> None:
|
||||
if isinstance(error, HTTPException):
|
||||
raise error
|
||||
def _reraise_api_error(error: Exception, path: str) -> NoReturn:
|
||||
if HTTPX_AVAILABLE and httpx is not None:
|
||||
if isinstance(error, httpx.TimeoutException):
|
||||
raise PeyeeyeGuardrailAPIError(f"peyeeye {path} timed out") from error
|
||||
|
|
|
|||
|
|
@ -97,9 +97,12 @@ async def test_pre_call_redacts_messages_and_caches_session():
|
|||
user = UserAPIKeyAuth(api_key="x")
|
||||
out = await g.async_pre_call_hook(user, cache, data, "completion")
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.peyeeye.peyeeye import (
|
||||
global_cache,
|
||||
)
|
||||
assert out["messages"][0]["content"] == "hi [EMAIL_1]"
|
||||
cached = cache.get_cache("peyeeye_session:call-1")
|
||||
assert cached == "ses_abc"
|
||||
assert global_cache.get_cache("peyeeye_session:call-1") == "ses_abc"
|
||||
global_cache.delete_cache("peyeeye_session:call-1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -123,9 +126,13 @@ async def test_pre_call_stateless_returns_skey():
|
|||
}
|
||||
await g.async_pre_call_hook(UserAPIKeyAuth(api_key="x"), cache, data, "completion")
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.peyeeye.peyeeye import (
|
||||
global_cache,
|
||||
)
|
||||
sent_body = g.async_handler.post.call_args.kwargs["json"]
|
||||
assert sent_body["session"] == "stateless"
|
||||
assert cache.get_cache("peyeeye_session:call-2") == "skey_xyz"
|
||||
assert global_cache.get_cache("peyeeye_session:call-2") == "skey_xyz"
|
||||
global_cache.delete_cache("peyeeye_session:call-2")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -144,18 +151,29 @@ async def test_pre_call_skips_when_no_messages():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_rehydrates_response():
|
||||
async def test_pre_and_post_call_roundtrip_uses_shared_cache():
|
||||
"""End-to-end: pre-call seeds session id, post-call retrieves & rehydrates."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.peyeeye.peyeeye import (
|
||||
global_cache,
|
||||
)
|
||||
|
||||
g = PeyeeyeGuardrail(peyeeye_api_key="pk", guardrail_name="t")
|
||||
g.async_handler = MagicMock()
|
||||
g.async_handler.post = AsyncMock(
|
||||
return_value=_ok({"text": "Reply to alice@acme.com", "replaced": 1})
|
||||
side_effect=[
|
||||
_ok({"text": ["hi [EMAIL_1]"], "session_id": "ses_abc"}),
|
||||
_ok({"text": "Reply to alice@acme.com", "replaced": 1}),
|
||||
]
|
||||
)
|
||||
g.async_handler.delete = AsyncMock()
|
||||
|
||||
# Seed the session id into litellm.cache so the post-call hook finds it.
|
||||
litellm.cache = MagicMock()
|
||||
litellm.cache.get_cache = MagicMock(return_value="ses_abc")
|
||||
litellm.cache.delete_cache = MagicMock()
|
||||
cache = DualCache()
|
||||
data = {
|
||||
"messages": [{"role": "user", "content": "hi alice@acme.com"}],
|
||||
"litellm_call_id": "rt-1",
|
||||
}
|
||||
user = UserAPIKeyAuth(api_key="x")
|
||||
await g.async_pre_call_hook(user, cache, data, "completion")
|
||||
|
||||
response = litellm.ModelResponse()
|
||||
response.choices = [
|
||||
|
|
@ -167,15 +185,12 @@ async def test_post_call_rehydrates_response():
|
|||
),
|
||||
)
|
||||
]
|
||||
data = {"litellm_call_id": "call-1"}
|
||||
|
||||
out = await g.async_post_call_success_hook(
|
||||
data, UserAPIKeyAuth(api_key="x"), response
|
||||
{"litellm_call_id": "rt-1"}, user, response
|
||||
)
|
||||
assert out.choices[0].message.content == "Reply to alice@acme.com"
|
||||
g.async_handler.delete.assert_awaited()
|
||||
litellm.cache.delete_cache.assert_called_once_with("peyeeye_session:call-1")
|
||||
litellm.cache = None
|
||||
assert global_cache.get_cache("peyeeye_session:rt-1") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -184,7 +199,6 @@ async def test_post_call_noop_without_session():
|
|||
g.async_handler = MagicMock()
|
||||
g.async_handler.post = AsyncMock()
|
||||
|
||||
litellm.cache = None
|
||||
response = litellm.ModelResponse()
|
||||
response.choices = [
|
||||
litellm.utils.Choices(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue