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:
Tim 2026-04-25 22:06:53 -05:00
parent a636014380
commit 3eda6087d1
2 changed files with 39 additions and 25 deletions

View file

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

View file

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