mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
Merge 720dbc8548 into 22b36cbcf6
This commit is contained in:
commit
846b26f8bf
2 changed files with 89 additions and 3 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,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={
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue