mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(guardrails): add logging_only_scope to observe one direction
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
85dc7cb62e
commit
8d2355392a
6 changed files with 293 additions and 20 deletions
|
|
@ -21,6 +21,7 @@ from litellm.types.guardrails import (
|
|||
DynamicGuardrailParams,
|
||||
GuardrailEventHooks,
|
||||
LitellmParams,
|
||||
LoggingOnlyScope,
|
||||
Mode,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -175,6 +176,7 @@ class CustomGuardrail(CustomLogger):
|
|||
use_native_lifecycle_hooks: ClassVar[bool] = False
|
||||
|
||||
records_own_guardrail_information: ClassVar[bool] = False
|
||||
logging_only_scope: LoggingOnlyScope | None
|
||||
|
||||
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
|
||||
super().__init_subclass__(**kwargs)
|
||||
|
|
@ -246,6 +248,7 @@ class CustomGuardrail(CustomLogger):
|
|||
self.run_in_parallel: bool = run_in_parallel
|
||||
self.scan_raw_request: bool = scan_raw_request
|
||||
self.only_scan_new_messages: bool = only_scan_new_messages
|
||||
self.logging_only_scope = None
|
||||
|
||||
if supported_event_hooks:
|
||||
## validate event_hook is in supported_event_hooks
|
||||
|
|
@ -803,6 +806,13 @@ class CustomGuardrail(CustomLogger):
|
|||
def uses_apply_guardrail_interface(self) -> bool:
|
||||
return type(self).apply_guardrail is not CustomGuardrail.apply_guardrail
|
||||
|
||||
def supports_logging_only_scope(self) -> bool:
|
||||
return (
|
||||
self.uses_apply_guardrail_interface()
|
||||
and not self.use_native_lifecycle_hooks
|
||||
and type(self).async_logging_hook is CustomGuardrail.async_logging_hook
|
||||
)
|
||||
|
||||
def _deployment_hook_target(self) -> "CustomLogger":
|
||||
if not self.uses_apply_guardrail_interface() or self.use_native_lifecycle_hooks:
|
||||
return self
|
||||
|
|
@ -999,12 +1009,19 @@ class CustomGuardrail(CustomLogger):
|
|||
"litellm_call_id": kwargs.get("litellm_call_id"),
|
||||
"metadata": scratch_metadata,
|
||||
}
|
||||
await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self)
|
||||
if response is None:
|
||||
if self.logging_only_scope != "output":
|
||||
try:
|
||||
await translation.process_input_messages(data=scratch_request, guardrail_to_apply=self)
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e)
|
||||
if response is None or self.logging_only_scope == "input":
|
||||
return
|
||||
await output_translation.process_output_response(
|
||||
response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request
|
||||
)
|
||||
try:
|
||||
await output_translation.process_output_response(
|
||||
response=copy.deepcopy(response), guardrail_to_apply=self, request_data=scratch_request
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.warning("Guardrail %s: logging_only scan raised: %s", self.guardrail_name, e)
|
||||
|
||||
def supports_scan_only_tool_results(self) -> bool:
|
||||
"""Whether this guardrail can scan tool-result content.
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ from .guardrail_hooks.llm_as_a_judge import (
|
|||
initialize_guardrail as initialize_llm_as_a_judge,
|
||||
)
|
||||
from .guardrail_initializers import (
|
||||
_configured_event_hooks,
|
||||
initialize_bedrock,
|
||||
initialize_hide_secrets,
|
||||
initialize_lakera,
|
||||
|
|
@ -439,6 +440,20 @@ def _as_callback_tuple(
|
|||
def _configure_callback_scoping(
|
||||
custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams
|
||||
) -> None:
|
||||
custom_guardrail_callback.logging_only_scope = litellm_params.logging_only_scope
|
||||
logging_only_scope: Final = litellm_params.logging_only_scope
|
||||
if logging_only_scope is not None and GuardrailEventHooks.logging_only.value not in _configured_event_hooks(
|
||||
litellm_params.mode
|
||||
):
|
||||
raise ValueError(
|
||||
f"Guardrail {guardrail_name}: logging_only_scope is set, but mode does not include logging_only, "
|
||||
"so it would never apply. Add logging_only to mode or remove logging_only_scope."
|
||||
)
|
||||
if logging_only_scope in ("input", "output") and not custom_guardrail_callback.supports_logging_only_scope():
|
||||
raise ValueError(
|
||||
f"Guardrail {guardrail_name}: logging_only_scope={logging_only_scope!r} is not supported by this "
|
||||
"guardrail, whose logging_only hook scans on its own. Remove logging_only_scope."
|
||||
)
|
||||
for scoping_param in (
|
||||
"skip_system_message_in_guardrail",
|
||||
"skip_tool_message_in_guardrail",
|
||||
|
|
|
|||
|
|
@ -895,6 +895,8 @@ class ContentFilterConfigModel(BaseModel):
|
|||
|
||||
MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"})
|
||||
|
||||
LoggingOnlyScope = Literal["input", "output", "both"]
|
||||
|
||||
|
||||
class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch update guardrails
|
||||
api_key: str | None = Field(default=None, description="API key for the guardrail service")
|
||||
|
|
@ -1136,6 +1138,14 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
),
|
||||
)
|
||||
|
||||
logging_only_scope: LoggingOnlyScope | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"which direction a logging_only scan observes: 'input' (request), 'output' (response), or 'both' "
|
||||
"(default). Only applies to mode logging_only; pre_call/post_call on the same guardrail keep blocking."
|
||||
),
|
||||
)
|
||||
|
||||
@field_validator(
|
||||
"mode",
|
||||
"default_action",
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import signal
|
|||
import socket
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -1297,3 +1298,96 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway,
|
|||
for response in responses:
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.headers["content-type"].startswith("text/event-stream"), response.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("logging_only_scope", "scanned_directions"),
|
||||
(("input", ("request",)), ("output", ("response",)), ("both", ("request", "response"))),
|
||||
)
|
||||
def test_logging_only_scope_observes_only_the_configured_direction_without_blocking(
|
||||
gateway: Gateway, tmp_path: Path, logging_only_scope: str, scanned_directions: tuple[str, ...]
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
prompt: Final = "synthetic observed prompt " + identity
|
||||
reply: Final = "synthetic observed reply " + identity
|
||||
texts_by_direction: Final = {"request": [prompt], "response": [reply]}
|
||||
|
||||
def guardrail(request: Request) -> Reply:
|
||||
assert request.target == "/beta/litellm_basic_guardrail_api"
|
||||
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic observed denial"}).encode())
|
||||
|
||||
def provider(request: Request) -> Reply:
|
||||
assert request.target == "/v1/chat/completions"
|
||||
assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}]
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
with wire_server(guardrail) as policy, wire_server(provider) as upstream:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": identity,
|
||||
"litellm_params": {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "logging_only",
|
||||
"logging_only_scope": logging_only_scope,
|
||||
"default_on": True,
|
||||
"api_base": policy.url,
|
||||
"api_key": "synthetic-guardrail-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "logging-only-scope.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(api_base=upstream.url + "/v1")
|
||||
response: Final = candidate.request(
|
||||
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["choices"][0]["message"]["content"] == reply, response.text
|
||||
assert len(upstream.drain()) == 1
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
scans: Final = tuple(json.loads(scan.body) for scan in policy.drain())
|
||||
assert [(scan["input_type"], scan["texts"]) for scan in scans] == [
|
||||
(direction, texts_by_direction[direction]) for direction in scanned_directions
|
||||
], scans
|
||||
entries: Final = object_value(rows[0]["metadata"])["guardrail_information"]
|
||||
assert isinstance(entries, list), rows[0]
|
||||
assert [
|
||||
(entry["guardrail_name"], entry["guardrail_mode"], entry["guardrail_status"])
|
||||
for entry in map(object_value, entries)
|
||||
] == [(identity, "logging_only", "guardrail_intervened")] * len(scanned_directions), entries
|
||||
today: Final = datetime.now(timezone.utc).date().isoformat()
|
||||
guardrail_id: Final = next(
|
||||
object_value(row)["guardrail_id"]
|
||||
for row in candidate.get("/v2/guardrails/list")["guardrails"]
|
||||
if object_value(row)["guardrail_name"] == identity
|
||||
)
|
||||
detail: Final = eventually(
|
||||
lambda: candidate.request(
|
||||
"GET",
|
||||
f"/guardrails/usage/detail/{guardrail_id}",
|
||||
params={"start_date": today, "end_date": today},
|
||||
).json(),
|
||||
lambda body: body["requestsEvaluated"] >= len(scanned_directions),
|
||||
seconds=30,
|
||||
return_last_on_timeout=True,
|
||||
)
|
||||
assert detail["requestsEvaluated"] == len(scanned_directions), detail
|
||||
|
|
|
|||
|
|
@ -1,15 +1,18 @@
|
|||
from collections.abc import Iterable
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy.guardrails.guardrail_registry import (
|
||||
get_guardrail_initializer_from_hooks,
|
||||
GuardrailRegistry,
|
||||
InMemoryGuardrailHandler,
|
||||
get_guardrail_initializer_from_hooks,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Guardrail, LitellmParams
|
||||
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams, LoggingOnlyScope, Mode
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
||||
def test_get_guardrail_initializer_from_hooks():
|
||||
|
|
@ -932,6 +935,117 @@ class TestScanOnlyToolResultsInitRefusal:
|
|||
)
|
||||
|
||||
|
||||
class _LoggingOnlyScopeSupportedGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: str,
|
||||
logging_obj: object | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return inputs
|
||||
|
||||
|
||||
class _LoggingOnlyScopeUnsupportedGuardrail(_LoggingOnlyScopeSupportedGuardrail):
|
||||
async def async_logging_hook(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
result: object,
|
||||
call_type: str,
|
||||
) -> tuple[dict[str, object], object]:
|
||||
return kwargs, result
|
||||
|
||||
|
||||
class TestLoggingOnlyScopeValidation:
|
||||
def _initialize(
|
||||
self,
|
||||
mode: str | list[str] | Mode,
|
||||
scope: LoggingOnlyScope | None,
|
||||
callback_type: type[CustomGuardrail] = _LoggingOnlyScopeSupportedGuardrail,
|
||||
) -> CustomGuardrail:
|
||||
from litellm.proxy.guardrails import guardrail_registry as registry_module
|
||||
|
||||
guardrail_type: Final = "logging_only_scope_test"
|
||||
|
||||
def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail:
|
||||
return callback_type(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
event_hook=litellm_params.mode,
|
||||
)
|
||||
|
||||
registry_module.guardrail_initializer_registry[guardrail_type] = _initializer
|
||||
lists: Final = _all_callback_lists()
|
||||
snapshots: Final = [list(callback_list) for callback_list in lists]
|
||||
try:
|
||||
handler: Final = InMemoryGuardrailHandler()
|
||||
result: Final = handler.initialize_guardrail(
|
||||
guardrail={
|
||||
"guardrail_name": "logging-only-scope-guardrail",
|
||||
"litellm_params": {
|
||||
"guardrail": guardrail_type,
|
||||
"mode": mode,
|
||||
"logging_only_scope": scope,
|
||||
},
|
||||
}
|
||||
)
|
||||
assert result is not None
|
||||
callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
|
||||
assert callback is not None
|
||||
return callback
|
||||
finally:
|
||||
for callback_list, snapshot in zip(lists, snapshots):
|
||||
callback_list[:] = snapshot
|
||||
registry_module.guardrail_initializer_registry.pop(guardrail_type, None)
|
||||
|
||||
def test_scope_requires_logging_only_mode(self) -> None:
|
||||
with pytest.raises(ValueError, match="logging_only_scope is set") as exc_info:
|
||||
self._initialize(mode="pre_call", scope="input")
|
||||
|
||||
assert str(exc_info.value) == (
|
||||
"Guardrail logging-only-scope-guardrail: logging_only_scope is set, but mode does not include "
|
||||
"logging_only, so it would never apply. Add logging_only to mode or remove logging_only_scope."
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"mode",
|
||||
(
|
||||
"logging_only",
|
||||
["pre_call", "logging_only"],
|
||||
Mode(tags={"audit": "logging_only"}, default="pre_call"),
|
||||
),
|
||||
)
|
||||
def test_scope_accepts_logging_only_in_supported_mode_forms(self, mode: str | list[str] | Mode) -> None:
|
||||
callback: Final = self._initialize(mode=mode, scope="input")
|
||||
|
||||
assert callback.logging_only_scope == "input"
|
||||
|
||||
def test_directional_scope_rejected_when_guardrail_owns_logging_hook(self) -> None:
|
||||
with pytest.raises(ValueError, match="logging_only_scope='input' is not supported") as exc_info:
|
||||
self._initialize(
|
||||
mode="logging_only",
|
||||
scope="input",
|
||||
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
|
||||
)
|
||||
|
||||
assert str(exc_info.value) == (
|
||||
"Guardrail logging-only-scope-guardrail: logging_only_scope='input' is not supported by this "
|
||||
"guardrail, whose logging_only hook scans on its own. Remove logging_only_scope."
|
||||
)
|
||||
|
||||
def test_both_scope_accepted_when_guardrail_owns_logging_hook(self) -> None:
|
||||
callback: Final = self._initialize(
|
||||
mode="logging_only",
|
||||
scope="both",
|
||||
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
|
||||
)
|
||||
|
||||
assert callback.logging_only_scope == "both"
|
||||
|
||||
def test_invalid_scope_fails_litellm_params_validation(self) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
LitellmParams(guardrail="test", mode="logging_only", logging_only_scope="request")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_guardrail_in_db_raises_when_row_missing():
|
||||
prisma_client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ from litellm.integrations.custom_guardrail import (
|
|||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.guardrails import GuardrailEventHooks, LoggingOnlyScope, Mode
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
GenericGuardrailAPIInputs,
|
||||
|
|
@ -546,14 +546,10 @@ class TestApplyGuardrailCheck:
|
|||
class ParentGuardrail(CustomGuardrail):
|
||||
"""Parent that inherits apply_guardrail from CustomGuardrail"""
|
||||
|
||||
pass
|
||||
|
||||
# Child class that only inherits apply_guardrail (doesn't override)
|
||||
class ChildGuardrailWithoutOverride(ParentGuardrail):
|
||||
"""Child that only inherits apply_guardrail"""
|
||||
|
||||
pass
|
||||
|
||||
# Child class that overrides apply_guardrail
|
||||
class ChildGuardrailWithOverride(ParentGuardrail):
|
||||
"""Child that overrides apply_guardrail"""
|
||||
|
|
@ -2540,7 +2536,7 @@ class TestLoggingOnlyApplyGuardrail:
|
|||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
|
||||
ContentFilterGuardrail,
|
||||
)
|
||||
from litellm.types.guardrails import BlockedWord, ContentFilterAction, GuardrailEventHooks
|
||||
from litellm.types.guardrails import BlockedWord, ContentFilterAction
|
||||
|
||||
guardrail: Final = ContentFilterGuardrail(
|
||||
guardrail_name="content-review",
|
||||
|
|
@ -2577,6 +2573,32 @@ class TestLoggingOnlyApplyGuardrail:
|
|||
assert "standard_logging_guardrail_information" not in kwargs["litellm_params"]["metadata"]
|
||||
assert kwargs["standard_logging_object"] == {"guardrail_information": None}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"scope,expected_calls",
|
||||
(
|
||||
(None, [("request", ["hello there"]), ("response", ["general kenobi"])]),
|
||||
("both", [("request", ["hello there"]), ("response", ["general kenobi"])]),
|
||||
("input", [("request", ["hello there"])]),
|
||||
("output", [("response", ["general kenobi"])]),
|
||||
),
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_logging_only_scope_scans_configured_directions(
|
||||
self,
|
||||
scope: LoggingOnlyScope | None,
|
||||
expected_calls: list[tuple[str, list[str]]],
|
||||
) -> None:
|
||||
guardrail: Final = _ApplyOnlyObserver()
|
||||
guardrail.logging_only_scope = scope
|
||||
kwargs, response = _logged_call([{"role": "user", "content": "hello there"}])
|
||||
|
||||
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
|
||||
|
||||
assert guardrail.calls == expected_calls
|
||||
entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"]
|
||||
assert len(entries) == len(expected_calls)
|
||||
assert [entry["guardrail_mode"] for entry in entries] == ["logging_only"] * len(expected_calls)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_appends_to_pre_call_verdicts_without_duplicating_them(self):
|
||||
guardrail = _ApplyOnlyObserver()
|
||||
|
|
@ -2606,14 +2628,17 @@ class TestLoggingOnlyApplyGuardrail:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_verdict_is_recorded_without_raising(self):
|
||||
guardrail = _ApplyOnlyObserver(block=True)
|
||||
guardrail: Final = _ApplyOnlyObserver(block=True)
|
||||
kwargs, response = _logged_call([{"role": "user", "content": "flagged content"}])
|
||||
|
||||
out_kwargs, _ = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value)
|
||||
|
||||
assert guardrail.calls == [("request", ["flagged content"])]
|
||||
entries = out_kwargs["standard_logging_object"]["guardrail_information"]
|
||||
assert [e["guardrail_status"] for e in entries] == ["guardrail_intervened"]
|
||||
assert guardrail.calls == [
|
||||
("request", ["flagged content"]),
|
||||
("response", ["general kenobi"]),
|
||||
]
|
||||
entries: Final = out_kwargs["standard_logging_object"]["guardrail_information"]
|
||||
assert [entry["guardrail_status"] for entry in entries] == ["guardrail_intervened", "guardrail_intervened"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_type_without_translation_is_skipped(self):
|
||||
|
|
@ -2939,9 +2964,7 @@ async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response(
|
|||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail = _NativeLifecycleLoggingGuardrail()
|
||||
assembled = ModelResponse(
|
||||
choices=[Choices(message=Message(role="assistant", content="assembled stream text"))]
|
||||
)
|
||||
assembled = ModelResponse(choices=[Choices(message=Message(role="assistant", content="assembled stream text"))])
|
||||
sentinel_result = object()
|
||||
kwargs = {
|
||||
"model": "gpt-5.4-mini",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue