mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(aporia): inject the HTTP handler and drop the explanatory comments
The tests replaced the guardrail's async_handler attribute after construction. The constructor now takes an optional async_handler, and the tests pass a mocked one in The verdict check loses its comment block, and the test helper takes a typed response body instead of Any
This commit is contained in:
parent
a708475f28
commit
3105fded3f
2 changed files with 38 additions and 53 deletions
|
|
@ -24,6 +24,7 @@ from litellm.litellm_core_utils.logging_utils import (
|
|||
convert_litellm_response_object_to_str,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
|
|
@ -44,9 +45,17 @@ class AporiaGuardrail(CustomGuardrail):
|
|||
GuardrailEventHooks.post_call,
|
||||
]
|
||||
|
||||
def __init__(self, api_key: str | None = None, api_base: str | None = None, **kwargs):
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
async_handler: AsyncHTTPHandler | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
|
||||
self.async_handler = async_handler or get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback
|
||||
)
|
||||
self.aporia_api_key = api_key or os.environ["APORIO_API_KEY"]
|
||||
self.aporia_api_base = api_base or os.environ["APORIO_API_BASE"]
|
||||
super().__init__(**kwargs)
|
||||
|
|
@ -129,13 +138,6 @@ class AporiaGuardrail(CustomGuardrail):
|
|||
# check if the response was flagged
|
||||
_json_response: Final = response.json()
|
||||
action: str = _json_response.get("action") # possible values are modify, passthrough, block, rephrase
|
||||
# Only `passthrough` says the content may go out as written.
|
||||
# `modify` and `rephrase` are interventions: Aporia is saying this
|
||||
# should not ship unchanged, and litellm cannot apply either from
|
||||
# this response, so forwarding the original would defeat the
|
||||
# guardrail while reporting success. An unrecognised or missing
|
||||
# action is treated the same way: a guardrail that cannot tell what
|
||||
# it was told must not be the one to decide the content is fine.
|
||||
if action != "passthrough":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
|
|||
|
|
@ -1,36 +1,28 @@
|
|||
"""Tests for how the Aporia guardrail acts on Aporia's verdict.
|
||||
|
||||
Aporia answers with one of four actions. Only ``passthrough`` says the content
|
||||
may go out as written; ``modify`` and ``rephrase`` are interventions. Forwarding
|
||||
the original for those defeats the guardrail and reports success while doing it
|
||||
(#41097).
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy.guardrails.guardrail_hooks.aporia_ai.aporia_ai import AporiaGuardrail
|
||||
|
||||
|
||||
def _guardrail(action: Any, include_action: bool = True) -> AporiaGuardrail:
|
||||
"""An AporiaGuardrail whose /validate call answers with one verdict."""
|
||||
guardrail = AporiaGuardrail(
|
||||
guardrail_name="aporia",
|
||||
api_key="k",
|
||||
api_base="https://example.invalid",
|
||||
)
|
||||
body = {"action": action} if include_action else {"reason": "no action key"}
|
||||
def _guardrail(body: dict[str, str | None]) -> tuple[AporiaGuardrail, AsyncMock]:
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.text = json.dumps(body)
|
||||
response.json = MagicMock(return_value=body)
|
||||
guardrail.async_handler = MagicMock()
|
||||
guardrail.async_handler.post = AsyncMock(return_value=response)
|
||||
return guardrail
|
||||
post = AsyncMock(return_value=response)
|
||||
handler = MagicMock(spec=AsyncHTTPHandler)
|
||||
handler.post = post
|
||||
guardrail = AporiaGuardrail(
|
||||
guardrail_name="aporia",
|
||||
api_key="k",
|
||||
api_base="https://example.invalid",
|
||||
async_handler=handler,
|
||||
)
|
||||
return guardrail, post
|
||||
|
||||
|
||||
async def _validate(guardrail: AporiaGuardrail) -> None:
|
||||
|
|
@ -42,21 +34,19 @@ async def _validate(guardrail: AporiaGuardrail) -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passthrough_is_forwarded():
|
||||
"""The accept control: a cleanly scanned request must still go out, or the
|
||||
guardrail is a wall rather than a filter."""
|
||||
guardrail = _guardrail("passthrough")
|
||||
guardrail, post = _guardrail({"action": "passthrough"})
|
||||
|
||||
await _validate(guardrail)
|
||||
|
||||
# It reached Aporia and came back without raising, which is the whole
|
||||
# observable effect of letting a request through.
|
||||
guardrail.async_handler.post.assert_awaited_once()
|
||||
post.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_is_refused():
|
||||
guardrail, _ = _guardrail({"action": "block"})
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _validate(_guardrail("block"))
|
||||
await _validate(guardrail)
|
||||
|
||||
assert exc.value.status_code == 400
|
||||
assert "Violated guardrail policy" in str(exc.value.detail)
|
||||
|
|
@ -64,11 +54,11 @@ async def test_block_is_refused():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("action", ["modify", "rephrase"])
|
||||
async def test_an_intervention_is_not_forwarded_unchanged(action):
|
||||
"""Aporia is saying this should not ship as written, and litellm cannot
|
||||
apply the modification from this response — so it must not ship it."""
|
||||
async def test_an_intervention_is_not_forwarded_unchanged(action: str):
|
||||
guardrail, _ = _guardrail({"action": action})
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await _validate(_guardrail(action))
|
||||
await _validate(guardrail)
|
||||
|
||||
assert exc.value.status_code == 400
|
||||
assert action in str(exc.value.detail)
|
||||
|
|
@ -76,18 +66,11 @@ async def test_an_intervention_is_not_forwarded_unchanged(action):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"action,include_action",
|
||||
[
|
||||
("BLOCK", True),
|
||||
("blocked", True),
|
||||
("", True),
|
||||
(None, True),
|
||||
(None, False),
|
||||
],
|
||||
"body",
|
||||
[{"action": "BLOCK"}, {"action": "blocked"}, {"action": ""}, {"action": None}, {}],
|
||||
)
|
||||
async def test_a_verdict_it_cannot_read_is_not_taken_for_permission(action, include_action):
|
||||
"""A response shape change, or an error object with no action at all, used
|
||||
to forward. A guardrail that cannot tell what it was told is not the one to
|
||||
decide the content is fine."""
|
||||
async def test_a_verdict_it_cannot_read_is_not_taken_for_permission(body: dict[str, str | None]):
|
||||
guardrail, _ = _guardrail(body)
|
||||
|
||||
with pytest.raises(HTTPException):
|
||||
await _validate(_guardrail(action, include_action=include_action))
|
||||
await _validate(guardrail)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue