diff --git a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py index dafa6e06652..21211da5fa5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py @@ -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={ diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_aporia.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_aporia.py new file mode 100644 index 00000000000..86077db7e4e --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_aporia.py @@ -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)