This commit is contained in:
Lukas 2026-09-27 16:20:51 -04:00 • committed by GitHub
commit 846b26f8bf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 89 additions and 3 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,7 +138,7 @@ 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
if action == "block":
if action != "passthrough":
raise HTTPException(
status_code=400,
detail={

View file

@ -0,0 +1,77 @@
import json
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(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)
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:
await guardrail.make_aporia_api_request(
request_data={},
new_messages=[{"role": "user", "content": "hello"}],
)
@pytest.mark.asyncio
async def test_passthrough_is_forwarded():
guardrail, post = _guardrail({"action": "passthrough"})
await _validate(guardrail)
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)
assert exc.value.status_code == 400
assert "Violated guardrail policy" in str(exc.value.detail)
@pytest.mark.asyncio
@pytest.mark.parametrize("action", ["modify", "rephrase"])
async def test_an_intervention_is_not_forwarded_unchanged(action: str):
guardrail, _ = _guardrail({"action": action})
with pytest.raises(HTTPException) as exc:
await _validate(guardrail)
assert exc.value.status_code == 400
assert action in str(exc.value.detail)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"body",
[{"action": "BLOCK"}, {"action": "blocked"}, {"action": ""}, {"action": None}, {}],
ids=["BLOCK", "blocked", "empty", "null", "missing"],
)
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)