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:
L4XB 2026-09-23 20:37:51 +02:00
parent a708475f28
commit 3105fded3f
No known key found for this signature in database
2 changed files with 38 additions and 53 deletions

View file

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

View file

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