mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(guardrails): validate logging-only mode with production integration coverage
This commit is contained in:
parent
02572d815f
commit
a459ac932f
3 changed files with 198 additions and 4 deletions
|
|
@ -601,6 +601,12 @@ class CustomGuardrail(CustomLogger):
|
|||
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None,
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
) -> None:
|
||||
allowed_hooks: Final = frozenset(supported_event_hooks) | (
|
||||
frozenset((GuardrailEventHooks.logging_only,))
|
||||
if self.uses_apply_guardrail_interface() and not self.use_native_lifecycle_hooks
|
||||
else frozenset()
|
||||
)
|
||||
|
||||
def _validate_event_hook_list_is_in_supported_event_hooks(
|
||||
event_hook: list[GuardrailEventHooks] | list[str],
|
||||
supported_event_hooks: list[GuardrailEventHooks],
|
||||
|
|
@ -608,7 +614,7 @@ class CustomGuardrail(CustomLogger):
|
|||
for hook in event_hook:
|
||||
if isinstance(hook, str):
|
||||
hook = GuardrailEventHooks(hook)
|
||||
if hook not in supported_event_hooks:
|
||||
if hook not in allowed_hooks:
|
||||
raise ValueError(f"Event hook {hook} is not in the supported event hooks {supported_event_hooks}")
|
||||
|
||||
if event_hook is None:
|
||||
|
|
@ -629,7 +635,7 @@ class CustomGuardrail(CustomLogger):
|
|||
default_list = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default]
|
||||
_validate_event_hook_list_is_in_supported_event_hooks(default_list, supported_event_hooks)
|
||||
elif isinstance(event_hook, GuardrailEventHooks):
|
||||
if event_hook not in supported_event_hooks:
|
||||
if event_hook not in allowed_hooks:
|
||||
raise ValueError(f"Event hook {event_hook} is not in the supported event hooks {supported_event_hooks}")
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import asyncio
|
||||
from typing import TYPE_CHECKING, Literal, Optional
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -2383,6 +2383,30 @@ class TestLoggingOnlyApplyGuardrail:
|
|||
apply_guardrail must still run against the logged request and response and
|
||||
record guardrail_information, instead of inheriting the CustomLogger no-op."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_content_filter_accepts_logging_only_and_records_detection(self):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
|
||||
ContentFilterGuardrail,
|
||||
)
|
||||
from litellm.types.guardrails import BlockedWord, ContentFilterAction, GuardrailEventHooks
|
||||
|
||||
guardrail: Final = ContentFilterGuardrail(
|
||||
guardrail_name="content-review",
|
||||
event_hook=GuardrailEventHooks.logging_only,
|
||||
default_on=True,
|
||||
blocked_words=[BlockedWord(keyword="hello", action=ContentFilterAction.BLOCK)],
|
||||
)
|
||||
kwargs, response = _logged_call([{"role": "user", "content": "hello there"}])
|
||||
|
||||
out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
|
||||
|
||||
assert out_response is response
|
||||
assert out_kwargs["messages"] == kwargs["messages"]
|
||||
assert (
|
||||
out_kwargs["standard_logging_object"]["guardrail_information"][0]["guardrail_status"]
|
||||
== "guardrail_intervened"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_apply_guardrail_observe_only_and_records_verdict(self):
|
||||
guardrail = _ApplyOnlyObserver()
|
||||
|
|
|
|||
|
|
@ -2,12 +2,17 @@ import json
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from opentelemetry.trace import StatusCode
|
||||
from prometheus_client import CollectorRegistry, Counter
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.generic_api.generic_api_callback import GenericAPILogger
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import GenericGuardrailAPI
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import ContentFilterGuardrail
|
||||
from litellm.types.guardrails import BlockedWord, ContentFilterAction, GuardrailEventHooks
|
||||
from litellm.types.utils import CallTypes
|
||||
from tests._prometheus_helpers import isolated_prometheus_registry
|
||||
from tests.test_litellm_rust.callback_recorder import RecordingLogger, drain_logging
|
||||
|
|
@ -32,6 +37,165 @@ pytestmark = pytest.mark.requires_rust_extension
|
|||
FAILURE_RESPONSE: Final = ResponseSpec(body={"message": "provider unavailable"}, status=500)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route", ASYNC_ROUTES, ids=route_id)
|
||||
async def test_generic_api_logger_exports_success_over_http(route: Route, provider: RecordingServer) -> None:
|
||||
provider.expected_requests = 2
|
||||
logger: Final = GenericAPILogger(endpoint=f"{provider.base_url}/logs", batch_size=1, log_format="single")
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
await route.invoke(provider, callbacks=[logger, recorder])
|
||||
await recorder.wait_for_async("async_log_success_event")
|
||||
|
||||
exports: Final = [request for request in provider.requests if request.path == "/logs"]
|
||||
assert len(exports) == 1
|
||||
payload: Final = exports[0].body
|
||||
assert payload["status"] == "success"
|
||||
assert payload["call_type"] == route.call_type
|
||||
assert payload["model"] == route.provider_model
|
||||
assert payload["response_cost"] == pytest.approx(route.expected_cost)
|
||||
assert route.response_text in json.dumps(payload["response"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"route",
|
||||
(
|
||||
OCR_ASYNC,
|
||||
pytest.param(
|
||||
MESSAGES_ROUTE,
|
||||
marks=pytest.mark.xfail(
|
||||
strict=True,
|
||||
raises=pytest.fail.Exception,
|
||||
reason="Native Messages provider failure raises TypeError and falls back to a second Python request",
|
||||
),
|
||||
),
|
||||
),
|
||||
ids=route_id,
|
||||
)
|
||||
async def test_generic_api_logger_exports_provider_failure_over_http(route: Route, provider: RecordingServer) -> None:
|
||||
provider.expected_requests = None
|
||||
provider.enqueue(FAILURE_RESPONSE)
|
||||
logger: Final = GenericAPILogger(endpoint=f"{provider.base_url}/logs", batch_size=1, log_format="single")
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
try:
|
||||
with pytest.raises(litellm.InternalServerError):
|
||||
await route.invoke(provider, callbacks=[logger, recorder], num_retries=0)
|
||||
finally:
|
||||
await drain_logging()
|
||||
await recorder.wait_for_async("async_log_failure_event")
|
||||
|
||||
exports: Final = [request for request in provider.requests if request.path == "/logs"]
|
||||
assert len(provider.requests) == 2
|
||||
assert len(exports) == 1
|
||||
payload: Final = exports[0].body
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["call_type"] == route.call_type
|
||||
assert payload["error_information"]["error_class"] == "InternalServerError"
|
||||
assert payload["error_information"]["error_code"] == "500"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route", NON_STREAM_ASYNC_ROUTES, ids=route_id)
|
||||
@pytest.mark.parametrize(
|
||||
("verdict", "expected_status"),
|
||||
(
|
||||
(ResponseSpec(body={"action": "NONE"}), "success"),
|
||||
(ResponseSpec(body={"action": "BLOCKED", "blocked_reason": "policy violation"}), "guardrail_intervened"),
|
||||
(ResponseSpec(body={"action": "GUARDRAIL_INTERVENED", "texts": ["redacted"]}), "success"),
|
||||
(FAILURE_RESPONSE, "guardrail_failed_to_respond"),
|
||||
),
|
||||
ids=("allow", "block", "rewrite", "unavailable"),
|
||||
)
|
||||
async def test_generic_guardrail_logging_only_verdict_is_exported_over_http(
|
||||
route: Route, provider: RecordingServer, verdict: ResponseSpec, expected_status: str
|
||||
) -> None:
|
||||
provider.expected_requests = 3
|
||||
provider.enqueue(ResponseSpec(body=route.provider_response))
|
||||
provider.enqueue(verdict)
|
||||
provider.enqueue(ResponseSpec(body={}))
|
||||
guardrail: Final = GenericGuardrailAPI(
|
||||
api_base=provider.base_url,
|
||||
guardrail_name="http-review",
|
||||
event_hook=GuardrailEventHooks.logging_only,
|
||||
default_on=True,
|
||||
)
|
||||
logger: Final = GenericAPILogger(endpoint=f"{provider.base_url}/logs", batch_size=1, log_format="single")
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
response: Final = await route.invoke(provider, callbacks=[guardrail, logger, recorder])
|
||||
await recorder.wait_for_async("async_log_success_event")
|
||||
|
||||
scan: Final = provider.requests[1]
|
||||
assert scan.path == "/beta/litellm_basic_guardrail_api"
|
||||
assert scan.body["input_type"] == route.logging_only_scan[0]
|
||||
assert scan.body["texts"] == list(route.logging_only_scan[1])
|
||||
assert provider.requests[2].path == "/logs"
|
||||
payload: Final = provider.requests[2].body
|
||||
assert payload["status"] == "success"
|
||||
assert [(entry["guardrail_name"], entry["guardrail_status"]) for entry in payload["guardrail_information"]] == [
|
||||
("http-review", expected_status)
|
||||
]
|
||||
assert route.response_text in json.dumps(payload["response"])
|
||||
if route == OCR_ASYNC:
|
||||
assert response.pages[0].markdown == route.response_text
|
||||
else:
|
||||
assert response["content"][0]["text"] == route.response_text
|
||||
assert "guardrails" not in provider.requests[0].body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route", NON_STREAM_ASYNC_ROUTES, ids=route_id)
|
||||
@pytest.mark.parametrize("action", (ContentFilterAction.BLOCK, ContentFilterAction.MASK))
|
||||
async def test_content_filter_logging_only_detects_real_content_without_changing_response(
|
||||
route: Route, provider: RecordingServer, action: ContentFilterAction
|
||||
) -> None:
|
||||
guardrail: Final = ContentFilterGuardrail(
|
||||
guardrail_name="content-review",
|
||||
event_hook=GuardrailEventHooks.logging_only,
|
||||
default_on=True,
|
||||
blocked_words=[BlockedWord(keyword=route.logging_only_scan[1][0], action=action)],
|
||||
)
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
response: Final = await route.invoke(provider, callbacks=[guardrail, recorder])
|
||||
payload: Final = (await recorder.wait_for_async("async_log_success_event"))[0].kwargs["standard_logging_object"]
|
||||
|
||||
assert len(payload["guardrail_information"]) == 1
|
||||
verdict: Final = payload["guardrail_information"][0]
|
||||
assert verdict["guardrail_name"] == "content-review"
|
||||
assert verdict["guardrail_status"] == ("guardrail_intervened" if action == ContentFilterAction.BLOCK else "success")
|
||||
assert verdict["guardrail_response"] == [
|
||||
{"action": action.value, "keyword": route.logging_only_scan[1][0].lower(), "type": "blocked_word"}
|
||||
]
|
||||
assert route.response_text in json.dumps(payload["response"])
|
||||
if route == OCR_ASYNC:
|
||||
assert response.pages[0].markdown == route.response_text
|
||||
else:
|
||||
assert response["content"][0]["text"] == route.response_text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("route", NON_STREAM_ASYNC_ROUTES, ids=route_id)
|
||||
@pytest.mark.xfail(
|
||||
strict=True,
|
||||
raises=pytest.fail.Exception,
|
||||
reason="CustomGuardrail deployment post-call hook does not dispatch to production apply_guardrail implementations",
|
||||
)
|
||||
async def test_content_filter_post_call_blocks_provider_response(route: Route, provider: RecordingServer) -> None:
|
||||
guardrail: Final = ContentFilterGuardrail(
|
||||
guardrail_name="enforced-content-review",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
blocked_words=[BlockedWord(keyword=route.response_text, action=ContentFilterAction.BLOCK)],
|
||||
)
|
||||
litellm.callbacks.append(guardrail)
|
||||
|
||||
with pytest.raises(HTTPException, match="Content blocked") as blocked:
|
||||
await route.invoke(provider, guardrails=["enforced-content-review"])
|
||||
assert blocked.value.status_code == 400
|
||||
|
||||
|
||||
def test_prometheus_registry_restores_collectors_after_failure() -> None:
|
||||
registry: Final = CollectorRegistry()
|
||||
original: Final = Counter("original", "Original collector", registry=registry)
|
||||
|
|
@ -174,7 +338,7 @@ async def test_prometheus_counts_one_failed_request(
|
|||
@pytest.mark.parametrize("route", NON_STREAM_ASYNC_ROUTES, ids=route_id)
|
||||
async def test_prometheus_by_string_name_is_initialized_once(route: Route, provider: RecordingServer) -> None:
|
||||
provider.expected_requests = 2
|
||||
litellm.success_callback = ["prometheus"] # test-quality-ok: tests public string-name registration; isolate_rust_state restores this registry
|
||||
litellm.success_callback = ["prometheus"] # test-quality-ok: public registration; fixture restores globals
|
||||
recorder: Final = RecordingLogger()
|
||||
|
||||
await route.invoke(provider, callbacks=[recorder])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue