From d011f18d65abc5e5fb4050a16103de19cf6063ef Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 15:19:58 -0400 Subject: [PATCH] fix(guardrails): address TrendAI review and CI lint findings --- .../guardrail_hooks/trendai/__init__.py | 1 - .../guardrail_hooks/trendai/_models.py | 4 +-- .../guardrail_hooks/trendai/trendai.py | 20 ++++++------ .../guardrail_hooks/trendai/test_trendai.py | 32 ++++++++----------- 4 files changed, 25 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py index 7803ac0baea..27b2d1206f7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py index 0acaf1fb0d3..b3ece6e2973 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py @@ -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") diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py index c6753738ffc..c55dfe1df58 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -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: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py index 4ae752cecd4..a2b197da315 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py @@ -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",