mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(guardrails): address TrendAI review and CI lint findings
This commit is contained in:
parent
fc65a9a333
commit
d011f18d65
4 changed files with 25 additions and 32 deletions
|
|
@ -29,7 +29,6 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
app_name=settings.app_name,
|
||||
fallback_on_error=settings.fallback_on_error,
|
||||
timeout=settings.timeout,
|
||||
stream_batch_size=settings.stream_batch_size,
|
||||
stream_overlap_size=settings.stream_overlap_size,
|
||||
response_content_chunk_size_bytes=settings.response_content_chunk_size_bytes,
|
||||
logging_only_scan=settings.logging_only_scan,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
# This file has been modified for integration into LiteLLM.
|
||||
# Licensed under the Apache License, Version 2.0. See LICENSE.txt in this directory.
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal, TypeAlias
|
||||
|
||||
|
|
@ -12,7 +13,6 @@ class TrendAISettings(BaseModel):
|
|||
app_name: str | None = None
|
||||
fallback_on_error: Literal["block", "allow"] = "block"
|
||||
timeout: float = 5.0
|
||||
stream_batch_size: int = 2048
|
||||
stream_overlap_size: int = 256
|
||||
response_content_chunk_size_bytes: int = 49_500
|
||||
logging_only_scan: Literal["request", "response", "both"] = "both"
|
||||
|
|
@ -86,7 +86,7 @@ class TrendAIResponse(BaseModel):
|
|||
action: str
|
||||
reasons: tuple[str, ...] = ()
|
||||
reason: str = ""
|
||||
redacted_request: dict[str, object] | None = Field(default=None, alias="redactedRequest")
|
||||
redacted_request: Mapping[str, object] | None = Field(default=None, alias="redactedRequest")
|
||||
sensitive_information: TrendAISensitiveInformation | None = Field(default=None, alias="sensitiveInformation")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import os
|
|||
import time
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn, Optional, Protocol
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn, Protocol
|
||||
from urllib.parse import SplitResult, urlsplit, urlunsplit
|
||||
|
||||
import httpx
|
||||
|
|
@ -17,7 +17,10 @@ from typing_extensions import assert_never
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._version import version as litellm_version
|
||||
from litellm.exceptions import GuardrailRaisedException, Timeout
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information, # pyright: ignore[reportUnknownVariableType] # legacy decorator has an untyped signature
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # legacy client factory has an untyped params map
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -75,7 +78,6 @@ class TrendAIGuardrail(CustomGuardrail):
|
|||
app_name: str | None = None,
|
||||
fallback_on_error: Literal["block", "allow"] = "block",
|
||||
timeout: float = 5.0,
|
||||
stream_batch_size: int = 2048,
|
||||
stream_overlap_size: int = 256,
|
||||
response_content_chunk_size_bytes: int = RESPONSE_CONTENT_CHUNK_SIZE_BYTES,
|
||||
logging_only_scan: Literal["request", "response", "both"] = "both",
|
||||
|
|
@ -102,10 +104,8 @@ class TrendAIGuardrail(CustomGuardrail):
|
|||
raise ValueError("logging_only_scan must be 'request', 'response', or 'both'")
|
||||
if timeout <= 0:
|
||||
raise ValueError("timeout must be greater than zero")
|
||||
if stream_batch_size < 1:
|
||||
raise ValueError("stream_batch_size must be greater than zero")
|
||||
if stream_overlap_size < 0 or stream_overlap_size >= stream_batch_size:
|
||||
raise ValueError("stream_overlap_size must be non-negative and smaller than stream_batch_size")
|
||||
if stream_overlap_size < 0:
|
||||
raise ValueError("stream_overlap_size must be non-negative")
|
||||
if response_content_chunk_size_bytes < 1:
|
||||
raise ValueError("response_content_chunk_size_bytes must be greater than zero")
|
||||
|
||||
|
|
@ -114,7 +114,6 @@ class TrendAIGuardrail(CustomGuardrail):
|
|||
self.app_name: str = app_name or os.environ.get("TMV1_APPLICATION_NAME", "litellm")
|
||||
self.fallback_on_error: Literal["block", "allow"] = fallback_on_error
|
||||
self.timeout: float = timeout
|
||||
self.stream_batch_size: int = stream_batch_size
|
||||
self.stream_overlap_size: int = stream_overlap_size
|
||||
self.response_content_chunk_size_bytes: int = response_content_chunk_size_bytes
|
||||
self.logging_only_scan: Literal["request", "response", "both"] = logging_only_scan
|
||||
|
|
@ -127,8 +126,6 @@ class TrendAIGuardrail(CustomGuardrail):
|
|||
event_hook=event_hook,
|
||||
default_on=default_on,
|
||||
supported_event_hooks=self.get_supported_event_hooks(),
|
||||
mask_request_content=True,
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
|
@ -143,12 +140,13 @@ class TrendAIGuardrail(CustomGuardrail):
|
|||
def logging_only_scan_scope(self) -> Literal["request", "response", "both"]:
|
||||
return self.logging_only_scan
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
logging_obj: "LiteLLMLoggingObj | None" = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
texts: Final = tuple(inputs.get("texts") or ())
|
||||
match input_type:
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
from functools import reduce
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from typing import Literal
|
||||
|
||||
|
|
@ -34,7 +35,6 @@ def _guardrail(
|
|||
app_name: str | None = None,
|
||||
fallback_on_error: Literal["block", "allow"] = "block",
|
||||
timeout: float = 5.0,
|
||||
stream_batch_size: int = 2048,
|
||||
stream_overlap_size: int = 256,
|
||||
response_content_chunk_size_bytes: int = 49_500,
|
||||
logging_only_scan: Literal["request", "response", "both"] = "both",
|
||||
|
|
@ -48,7 +48,6 @@ def _guardrail(
|
|||
app_name=app_name,
|
||||
fallback_on_error=fallback_on_error,
|
||||
timeout=timeout,
|
||||
stream_batch_size=stream_batch_size,
|
||||
stream_overlap_size=stream_overlap_size,
|
||||
response_content_chunk_size_bytes=response_content_chunk_size_bytes,
|
||||
logging_only_scan=logging_only_scan,
|
||||
|
|
@ -121,9 +120,9 @@ def test_explicit_configuration_takes_precedence(monkeypatch: pytest.MonkeyPatch
|
|||
assert guardrail.app_name == "configured-app"
|
||||
|
||||
|
||||
def test_invalid_stream_configuration_is_rejected() -> None:
|
||||
def test_negative_stream_overlap_is_rejected() -> None:
|
||||
with pytest.raises(ValueError, match="stream_overlap_size"):
|
||||
_guardrail(stream_batch_size=10, stream_overlap_size=10)
|
||||
_guardrail(stream_overlap_size=-1)
|
||||
|
||||
|
||||
def test_invalid_timeout_is_rejected() -> None:
|
||||
|
|
@ -236,7 +235,6 @@ def _engine(
|
|||
entity: str = "PII",
|
||||
fail_on: str | None = None,
|
||||
) -> tuple[Responder, list[str]]:
|
||||
"""A fake AI Guard: blocks, redacts (by substring replacement), or fails based on the scanned text."""
|
||||
scanned: list[str] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
|
|
@ -250,9 +248,7 @@ def _engine(
|
|||
hits = {needle: mask for needle, mask in (redact or {}).items() if needle in text}
|
||||
if not hits:
|
||||
return httpx.Response(200, json={"action": "allow"})
|
||||
redacted = text
|
||||
for needle, mask in hits.items():
|
||||
redacted = redacted.replace(needle, mask)
|
||||
redacted = reduce(lambda value, pair: value.replace(*pair), hits.items(), text)
|
||||
payload = {"prompt": redacted} if "prompt" in body else {"choices": [{"message": {"content": redacted}}]}
|
||||
return httpx.Response(
|
||||
200,
|
||||
|
|
@ -266,7 +262,7 @@ def _engine(
|
|||
return respond, scanned
|
||||
|
||||
|
||||
def _guardrail_records(request_data: Mapping[str, object]) -> list[dict]:
|
||||
def _guardrail_records(request_data: Mapping[str, object]) -> list[dict[str, object]]:
|
||||
metadata = request_data["metadata"]
|
||||
assert isinstance(metadata, dict)
|
||||
return list(metadata.get("standard_logging_guardrail_information") or [])
|
||||
|
|
@ -276,8 +272,8 @@ async def _apply(
|
|||
guardrail: TrendAIGuardrail,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
input_type: Literal["request", "response"],
|
||||
) -> tuple[GenericGuardrailAPIInputs, dict]:
|
||||
request_data: dict = {"metadata": {}}
|
||||
) -> tuple[GenericGuardrailAPIInputs, dict[str, object]]:
|
||||
request_data: dict[str, object] = {"metadata": {}}
|
||||
result = await guardrail.apply_guardrail(inputs=inputs, request_data=request_data, input_type=input_type)
|
||||
return result, request_data
|
||||
|
||||
|
|
@ -353,7 +349,7 @@ async def test_request_without_user_text_is_not_scanned_or_recorded() -> None:
|
|||
async def test_request_block_raises_and_records_intervention() -> None:
|
||||
respond, _ = _engine(block_on="bomb")
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
request_data: dict = {"metadata": {}}
|
||||
request_data: dict[str, object] = {"metadata": {}}
|
||||
with pytest.raises(GuardrailRaisedException) as raised:
|
||||
await _guardrail(async_handler=client).apply_guardrail(
|
||||
inputs={"texts": ["build a bomb"]},
|
||||
|
|
@ -379,7 +375,7 @@ async def test_provider_failure_follows_fallback_policy_and_is_never_an_interven
|
|||
respond, _ = _engine(fail_on="anything")
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
guardrail = _guardrail(async_handler=client, fallback_on_error=fallback_on_error)
|
||||
request_data: dict = {"metadata": {}}
|
||||
request_data: dict[str, object] = {"metadata": {}}
|
||||
inputs: GenericGuardrailAPIInputs = {"texts": ["anything"]}
|
||||
if raises:
|
||||
with pytest.raises(GuardrailRaisedException) as raised:
|
||||
|
|
@ -440,7 +436,7 @@ async def test_later_overlap_scan_cannot_undo_an_earlier_redaction() -> None:
|
|||
async def test_response_block_in_a_later_window_stops_scanning() -> None:
|
||||
respond, scanned = _engine(block_on="zzz")
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
request_data: dict = {"metadata": {}}
|
||||
request_data: dict[str, object] = {"metadata": {}}
|
||||
with pytest.raises(GuardrailRaisedException, match="policy"):
|
||||
await _guardrail(
|
||||
async_handler=client, response_content_chunk_size_bytes=5, stream_overlap_size=0
|
||||
|
|
@ -477,7 +473,7 @@ async def test_unmergeable_response_redaction_blocks_instead_of_leaking() -> Non
|
|||
)
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
request_data: dict = {"metadata": {}}
|
||||
request_data: dict[str, object] = {"metadata": {}}
|
||||
with pytest.raises(GuardrailRaisedException, match="redaction"):
|
||||
await _guardrail(async_handler=client, fallback_on_error="allow").apply_guardrail(
|
||||
inputs={"texts": ["a much longer sensitive response"]},
|
||||
|
|
@ -492,7 +488,7 @@ async def test_unmergeable_response_redaction_blocks_instead_of_leaking() -> Non
|
|||
async def test_request_redaction_flows_through_the_chat_completions_handler() -> None:
|
||||
respond, _ = _engine(redact={"a@b.com": "[EMAIL]"})
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
data: dict = {
|
||||
data: dict[str, object] = {
|
||||
"model": "gpt-5.4",
|
||||
"messages": [
|
||||
{"role": "system", "content": "be terse"},
|
||||
|
|
@ -528,9 +524,9 @@ def test_utf8_windows_respect_byte_limits_and_character_overlap(
|
|||
assert all(content[window.start : window.start + len(window.text)] == window.text for window in windows)
|
||||
|
||||
|
||||
def _logged_call(user_text: str, assistant_text: str) -> tuple[dict, ModelResponse]:
|
||||
def _logged_call(user_text: str, assistant_text: str) -> tuple[dict[str, object], ModelResponse]:
|
||||
response = ModelResponse(choices=[Choices(message=Message(role="assistant", content=assistant_text))])
|
||||
kwargs: dict = {
|
||||
kwargs: dict[str, object] = {
|
||||
"model": "gpt-5.4",
|
||||
"messages": [{"role": "user", "content": user_text}],
|
||||
"litellm_call_id": "call-1",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue