fix(guardrails): address TrendAI review and CI lint findings

This commit is contained in:
Oliver Fei 2026-09-24 15:19:58 -04:00
parent fc65a9a333
commit d011f18d65
4 changed files with 25 additions and 32 deletions

View file

@ -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,

View file

@ -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")

View file

@ -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:

View file

@ -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",