From 5701da08844b147cacebe8d4207034cce5f6fd32 Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Wed, 9 Sep 2026 14:25:38 -0400 Subject: [PATCH 1/7] feat(guardrails): add TrendAI Guard guardrail --- .../guardrail_hooks/trendai/LICENSE.txt | 201 +++++++ .../guardrail_hooks/trendai/__init__.py | 50 ++ .../guardrail_hooks/trendai/_models.py | 134 +++++ .../guardrail_hooks/trendai/_text.py | 157 +++++ .../guardrail_hooks/trendai/trendai.py | 513 ++++++++++++++++ .../guardrail_hooks/trendai/test_trendai.py | 563 ++++++++++++++++++ 6 files changed, 1618 insertions(+) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/trendai/LICENSE.txt create mode 100644 litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/LICENSE.txt b/litellm/proxy/guardrails/guardrail_hooks/trendai/LICENSE.txt new file mode 100644 index 00000000000..3d080f27b77 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/LICENSE.txt @@ -0,0 +1,201 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright [2023] [Trend Micro] + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py new file mode 100644 index 00000000000..b6ca0d3a7c7 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py @@ -0,0 +1,50 @@ +from typing import TYPE_CHECKING, Final + +import litellm +from litellm.types.guardrails import GuardrailEventHooks, Mode + +from ._models import TrendAISettings +from .trendai import GUARDRAIL_NAME, TrendAIGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def _normalize_event_hook( + mode: str | list[str] | Mode, +) -> GuardrailEventHooks | list[GuardrailEventHooks] | Mode: + if isinstance(mode, str): + return GuardrailEventHooks(mode) + if isinstance(mode, list): + return [GuardrailEventHooks(item) for item in mode] + return mode + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> TrendAIGuardrail: + settings: Final = TrendAISettings.model_validate(litellm_params.model_dump(mode="python")) + guardrail_name: Final = guardrail["guardrail_name"] + callback: Final = TrendAIGuardrail( + api_key=litellm_params.api_key, + api_base=litellm_params.api_base, + app_name=settings.app_name, + fallback_on_error=settings.fallback_on_error, + mask_pii=settings.mask_pii, + 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, + guardrail_name=guardrail_name, + event_hook=_normalize_event_hook(litellm_params.mode), + default_on=litellm_params.default_on is True, + ) + litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] # callback union is partially untyped + callback + ) + return callback + + +guardrail_initializer_registry: Final = {GUARDRAIL_NAME: initialize_guardrail} +guardrail_class_registry: Final = {GUARDRAIL_NAME: TrendAIGuardrail} + +__all__ = ("TrendAIGuardrail",) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py new file mode 100644 index 00000000000..2efcbb52c75 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py @@ -0,0 +1,134 @@ +# Derived from tm-v1-ai-guard-litellm-plugin revision 6cafc143f62962a98d4eb5abe9f608c61ff194d4. +# This file has been modified for integration into LiteLLM. +# Licensed under the Apache License, Version 2.0. See LICENSE.txt in this directory. + +from dataclasses import dataclass +from typing import Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field + + +class TrendAISettings(BaseModel): + app_name: str | None = None + fallback_on_error: Literal["block", "allow"] = "block" + mask_pii: bool = True + 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" + + +class TrendAIChatMessage(BaseModel): + role: Literal["assistant"] = "assistant" + content: str + + +class TrendAIChatChoice(BaseModel): + index: int = 0 + message: TrendAIChatMessage + finish_reason: Literal["stop"] = "stop" + + +class TrendAIChatCompletionPayload(BaseModel): + """The OpenAI chat-completion shape Trend AI Guard scans model output as.""" + + id: str = "chatcmpl-stream" + object: Literal["chat.completion"] = "chat.completion" + created: int = 0 + model: str + choices: tuple[TrendAIChatChoice, ...] + + @classmethod + def for_content(cls, content: str, model: str) -> "TrendAIChatCompletionPayload": + return cls(model=model, choices=(TrendAIChatChoice(message=TrendAIChatMessage(content=content)),)) + + +class TrendAIRedactedMessage(BaseModel): + model_config = ConfigDict(extra="ignore") + + content: str | None = None + + +class TrendAIRedactedChoice(BaseModel): + model_config = ConfigDict(extra="ignore") + + message: TrendAIRedactedMessage | None = None + + +class TrendAIRedactedResponse(BaseModel): + model_config = ConfigDict(extra="ignore") + + choices: tuple[TrendAIRedactedChoice, ...] = () + + +class TrendAIRedactedPrompt(BaseModel): + model_config = ConfigDict(extra="ignore") + + prompt: str | None = None + + +class TrendAISensitiveRule(BaseModel): + model_config = ConfigDict(extra="ignore") + + id: str = "" + + +class TrendAISensitiveInformation(BaseModel): + model_config = ConfigDict(extra="ignore") + + has_policy_violation: bool = Field(default=False, alias="hasPolicyViolation") + rules: tuple[TrendAISensitiveRule, ...] = () + + +class TrendAIResponse(BaseModel): + model_config = ConfigDict(extra="ignore") + + action: str + reasons: tuple[str, ...] = () + reason: str = "" + redacted_request: dict[str, object] | None = Field(default=None, alias="redactedRequest") + sensitive_information: TrendAISensitiveInformation | None = Field(default=None, alias="sensitiveInformation") + + +@dataclass(frozen=True, slots=True) +class TrendAIAllow: + kind: Literal["allow"] = "allow" + redacted_content: str | None = None + masked_entity_count: tuple[tuple[str, int], ...] = () + + +@dataclass(frozen=True, slots=True) +class TrendAIBlock: + reason: str + status_code: int = 400 + kind: Literal["block"] = "block" + + +@dataclass(frozen=True, slots=True) +class TrendAIProviderFailure: + reason: str + status_code: int | None = None + kind: Literal["provider_failure"] = "provider_failure" + + +TrendAIScanResult: TypeAlias = TrendAIAllow | TrendAIBlock | TrendAIProviderFailure + + +@dataclass(frozen=True, slots=True) +class TrendAITextWindow: + text: str + start: int + + +@dataclass(frozen=True, slots=True) +class TrendAIWindowScan: + window: TrendAITextWindow + result: TrendAIScanResult + + +@dataclass(frozen=True, slots=True) +class TrendAIRequestPrompt: + prompt: str + start: int + end: int diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py new file mode 100644 index 00000000000..4cf09afabb8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py @@ -0,0 +1,157 @@ +# Derived from tm-v1-ai-guard-litellm-plugin revision 6cafc143f62962a98d4eb5abe9f608c61ff194d4. +# 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 Iterator, Sequence +from typing import Final, Literal + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +from litellm.types.llms.openai import AllMessageValues + +from ._models import TrendAIRequestPrompt, TrendAITextWindow + + +class _TextPart(BaseModel): + model_config = ConfigDict(extra="ignore") + + type: Literal["text"] + text: str + + +class _OtherPart(BaseModel): + model_config = ConfigDict(extra="ignore") + + type: str + + +class _UserMessage(BaseModel): + model_config = ConfigDict(extra="ignore") + + role: str + content: str | tuple[_TextPart | _OtherPart, ...] | None = None + + +_MESSAGES: Final = TypeAdapter(tuple[_UserMessage, ...]) + + +def utf8_windows(content: str, *, chunk_size_bytes: int, overlap_chars: int) -> tuple[TrendAITextWindow, ...]: + """Split ``content`` into windows of at most ``chunk_size_bytes`` UTF-8 bytes. + + Consecutive windows share their last ``overlap_chars`` characters so a finding that straddles a + window boundary is still seen whole by one scan. The overlap is capped below the window length so + every window makes forward progress. + """ + return tuple(_iter_utf8_windows(content, chunk_size_bytes, overlap_chars)) + + +def _iter_utf8_windows(content: str, chunk_size_bytes: int, overlap_chars: int) -> Iterator[TrendAITextWindow]: + encoded: Final = content.encode("utf-8") + byte_offset = 0 # rebind-ok: window cursor advances across the loop + char_offset = 0 # rebind-ok: window cursor advances across the loop + while byte_offset < len(encoded): + text = ( + encoded[byte_offset : byte_offset + chunk_size_bytes].decode("utf-8", errors="ignore") + or content[char_offset] + ) + yield TrendAITextWindow(text=text, start=char_offset) + end_byte_offset = byte_offset + len(text.encode("utf-8")) + if end_byte_offset >= len(encoded): + return + overlap = min(max(overlap_chars, 0), len(text) - 1) + byte_offset = end_byte_offset - len(text[len(text) - overlap :].encode("utf-8")) + char_offset += len(text) - overlap + + +def merge_redaction(original: str, current: str, redacted: str) -> str | None: + """Apply the masks ``redacted`` adds over ``original`` on top of the masks already in ``current``. + + Returns None when the three texts do not line up character for character, since a positional + merge would then scramble the output. + """ + if len(original) != len(current) or len(redacted) != len(original): + return None + return "".join( + redacted_char if redacted_char != original_char else current_char + for original_char, current_char, redacted_char in zip(original, current, redacted, strict=True) + ) + + +def apply_window_redaction(content: str, window: TrendAITextWindow, redacted: str) -> str | None: + end: Final = window.start + len(window.text) + merged: Final = merge_redaction(window.text, content[window.start : end], redacted) + if merged is None: + return None + return f"{content[: window.start]}{merged}{content[end:]}" + + +def _text_parts(content: str | tuple[_TextPart | _OtherPart, ...] | None) -> tuple[str, ...] | None: + if content is None: + return None + if isinstance(content, str): + return (content,) + return tuple(part.text for part in content if isinstance(part, _TextPart)) + + +def _last_user_text_parts(structured_messages: Sequence[AllMessageValues]) -> tuple[str, ...] | None: + try: + messages: Final = _MESSAGES.validate_python(structured_messages) + except ValidationError: + return None + last_user_message: Final = next((message for message in reversed(messages) if message.role == "user"), None) + if last_user_message is None: + return None + return _text_parts(last_user_message.content) + + +def _last_occurrence(texts: Sequence[str], parts: Sequence[str]) -> int | None: + return next( + ( + start + for start in range(len(texts) - len(parts), -1, -1) + if tuple(texts[start : start + len(parts)]) == tuple(parts) + ), + None, + ) + + +def locate_request_prompt( + texts: Sequence[str], + structured_messages: Sequence[AllMessageValues] | None, +) -> TrendAIRequestPrompt | None: + """Pick the text Trend AI Guard scans for a request and where it lives in ``texts``. + + The scanned prompt is the last user turn, matching the upstream plugin. Its text parts are + located as a run inside ``texts`` so a redacted prompt can be written back to exactly those + entries. Without structured messages the whole ``texts`` list is the prompt. + """ + parts: Final = _last_user_text_parts(structured_messages) if structured_messages else tuple(texts) + if parts is None or not parts: + return None + prompt: Final = "".join(parts).strip() + if not prompt: + return None + start: Final = _last_occurrence(texts, parts) + if start is None: + return None + return TrendAIRequestPrompt(prompt=prompt, start=start, end=start + len(parts)) + + +def redact_request_texts(texts: Sequence[str], prompt: TrendAIRequestPrompt, redacted: str) -> tuple[str, ...]: + """Write a redacted prompt back over the ``texts`` entries the prompt was assembled from. + + A single-part prompt is replaced outright. A multi-part prompt whose redaction kept its length + is split back into the original parts; otherwise the first part carries the whole redaction and + the rest are blanked, so no original part can leak. + """ + parts: Final = tuple(texts[prompt.start : prompt.end]) + if len(parts) == 1: + return (*texts[: prompt.start], redacted, *texts[prompt.end :]) + joined: Final = "".join(parts) + leading: Final = len(joined) - len(joined.lstrip()) + aligned: Final = f"{joined[:leading]}{redacted}{joined[len(joined.rstrip()) :]}" + if len(aligned) != len(joined): + return (*texts[: prompt.start], redacted, *("" for _ in parts[1:]), *texts[prompt.end :]) + offsets: Final = tuple(sum(len(part) for part in parts[:index]) for index in range(len(parts) + 1)) + split: Final = tuple(aligned[offsets[index] : offsets[index + 1]] for index in range(len(parts))) + return (*texts[: prompt.start], *split, *texts[prompt.end :]) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py new file mode 100644 index 00000000000..4ed0e54abe8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -0,0 +1,513 @@ +# Derived from tm-v1-ai-guard-litellm-plugin revision 6cafc143f62962a98d4eb5abe9f608c61ff194d4. +# This file has been modified for integration into LiteLLM. +# Licensed under the Apache License, Version 2.0. See LICENSE.txt in this directory. + +import asyncio +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 urllib.parse import SplitResult, urlsplit, urlunsplit + +import httpx +from pydantic import ValidationError +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.llms.custom_httpx.http_handler import ( + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # legacy client factory has an untyped params map + httpxSpecialProvider, +) +from litellm.types.guardrails import GuardrailEventHooks, Mode +from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus + +from ._models import ( + TrendAIAllow, + TrendAIBlock, + TrendAIChatCompletionPayload, + TrendAIProviderFailure, + TrendAIRedactedPrompt, + TrendAIRedactedResponse, + TrendAIResponse, + TrendAIScanResult, + TrendAIWindowScan, +) +from ._text import apply_window_redaction, locate_request_prompt, redact_request_texts, utf8_windows + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +GUARDRAIL_NAME: Final = "trendai" +OPENAI_CHAT_COMPLETION_RESPONSE_V1: Final = "OpenAIChatCompletionResponseV1" +RESPONSE_CONTENT_CHUNK_SIZE_BYTES: Final = 49_500 +TMV1_CLIENT_NAME: Final = "litellm" +PLUGIN_VERSION: Final = "0.1.2" +PROVIDER_UNAVAILABLE_STATUS: Final = 503 +_APPLY_GUARDRAILS_PATH: Final = "/applyGuardrails" +_TREND_AI_SECURITY_PATH: Final = "/v3.0/aiSecurity" +_RESPONSE_MODEL: Final = "guardrailed-response" +_REDACTION_FAILED_MESSAGE: Final = "Trend AI Guard could not apply the requested redaction" +_NO_ENTITIES: Final[Mapping[str, int]] = MappingProxyType({}) + + +class _AsyncHTTPClient(Protocol): + async def post( + self, + url: str, + *, + json: dict[str, object] | None = None, + headers: dict[str, str] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> httpx.Response: ... + + +class TrendAIGuardrail(CustomGuardrail): + records_own_guardrail_information: ClassVar[bool] = True + + def __init__( + self, + api_key: str | None = None, + api_base: str | None = None, + app_name: str | None = None, + fallback_on_error: Literal["block", "allow"] = "block", + mask_pii: bool = True, + 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", + async_handler: _AsyncHTTPClient | None = None, + guardrail_name: str | None = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, + default_on: bool = False, + ) -> None: + resolved_api_key: Final = api_key or os.environ.get("TMV1_API_KEY") + if not resolved_api_key: + raise ValueError( + "Trend AI Guard requires an API key. Pass api_key or set the TMV1_API_KEY environment variable." + ) + + resolved_api_base: Final = api_base or os.environ.get("TRENDAI_AI_GUARD_BASE_URL") + if not resolved_api_base: + raise ValueError( + "Trend AI Guard requires an API base URL. Pass api_base or set the " + "TRENDAI_AI_GUARD_BASE_URL environment variable." + ) + if fallback_on_error not in ("block", "allow"): + raise ValueError("fallback_on_error must be 'block' or 'allow'") + if logging_only_scan not in ("request", "response", "both"): + 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 response_content_chunk_size_bytes < 1: + raise ValueError("response_content_chunk_size_bytes must be greater than zero") + + self.api_key: str = resolved_api_key + self.api_url: str = _build_apply_guardrails_url(resolved_api_base) + 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.mask_pii: bool = mask_pii + 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 + self.async_handler: _AsyncHTTPClient = async_handler or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + + super().__init__( # pyright: ignore[reportUnknownMemberType] # base constructor retains untyped extension kwargs + guardrail_name=guardrail_name, + event_hook=event_hook, + default_on=default_on, + supported_event_hooks=self.get_supported_event_hooks(), + mask_request_content=mask_pii, + mask_response_content=mask_pii, + ) + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: CustomGuardrail contract + return [ # mutable-ok: CustomGuardrail contract + GuardrailEventHooks.pre_call, + GuardrailEventHooks.during_call, + GuardrailEventHooks.post_call, + GuardrailEventHooks.logging_only, + ] + + def logging_only_scan_scope(self) -> Literal["request", "response", "both"]: + return self.logging_only_scan + + 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, + ) -> GenericGuardrailAPIInputs: + texts: Final = tuple(inputs.get("texts") or ()) + match input_type: + case "request": + return await self._apply_request_guardrail(inputs, texts, request_data) + case "response": + return await self._apply_response_guardrail(inputs, texts, request_data) + case _: + assert_never(input_type) + + async def _apply_request_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + texts: tuple[str, ...], + request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict + ) -> GenericGuardrailAPIInputs: + prompt: Final = locate_request_prompt(texts, inputs.get("structured_messages")) + if prompt is None: + verbose_proxy_logger.debug("Trend AI Guard: no user prompt to scan in request inputs") + return inputs + started_at: Final = time.time() + result: Final = await self._scan_payload(MappingProxyType({"prompt": prompt.prompt})) + self._record_scan(result, request_data, started_at, event_type=None) + redacted: Final = self._enforce(result) + if redacted is None: + return inputs + redacted_inputs: Final[GenericGuardrailAPIInputs] = { + **inputs, + "texts": list(redact_request_texts(texts, prompt, redacted)), # mutable-ok: GenericGuardrailAPIInputs field + } + return redacted_inputs + + async def _apply_response_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + texts: tuple[str, ...], + request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict + ) -> GenericGuardrailAPIInputs: + if not texts: + verbose_proxy_logger.debug("Trend AI Guard: no response text to scan") + return inputs + model: Final = inputs.get("model") or _RESPONSE_MODEL + started_at: Final = time.time() + scans: Final = tuple([await self._scan_response_windows(text, model) for text in texts]) + all_scans: Final = tuple(scan for text_scans in scans for scan in text_scans) + verdict: Final = _window_verdict(all_scans) + self._record_scan( + verdict, + request_data, + started_at, + event_type=None, + redacted=any(_is_redacting(scan.result) for scan in all_scans), + ) + self._enforce(verdict) + redacted_texts: Final = tuple( + _merge_window_redactions(text, text_scans) for text, text_scans in zip(texts, scans, strict=True) + ) + merged_texts: Final = tuple(text for text in redacted_texts if text is not None) + if len(merged_texts) != len(texts): + self._raise_redaction_failure(request_data, started_at, event_type=None) + redacted_inputs: Final[GenericGuardrailAPIInputs] = { + **inputs, + "texts": list(merged_texts), # mutable-ok: GenericGuardrailAPIInputs field + } + return redacted_inputs + + async def _scan_response_windows(self, content: str, model: str) -> tuple[TrendAIWindowScan, ...]: + return tuple([scan async for scan in self._iter_response_window_scans(content, model)]) + + async def _iter_response_window_scans(self, content: str, model: str) -> AsyncIterator[TrendAIWindowScan]: + """Scan ``content`` window by window, stopping at the first block or provider failure.""" + for window in utf8_windows( + content, + chunk_size_bytes=self.response_content_chunk_size_bytes, + overlap_chars=self.stream_overlap_size, + ): + result = await self._scan_payload( + _chat_completion_payload(window.text, model), + request_type=OPENAI_CHAT_COMPLETION_RESPONSE_V1, + ) + yield TrendAIWindowScan(window=window, result=result) + if not isinstance(result, TrendAIAllow): + return + + def _raise_redaction_failure( + self, + request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict + started_at: float, + *, + event_type: GuardrailEventHooks | None, + ) -> NoReturn: + self._record_failure(_REDACTION_FAILED_MESSAGE, request_data, started_at, event_type=event_type) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=_REDACTION_FAILED_MESSAGE, + should_wrap_with_default_message=False, + status_code=PROVIDER_UNAVAILABLE_STATUS, + ) + + def _enforce(self, result: TrendAIScanResult) -> str | None: + """Raise for a block or a fail-closed provider failure; otherwise return any redacted content.""" + match result: + case TrendAIAllow(redacted_content=redacted_content): + return redacted_content + case TrendAIBlock(reason=reason, status_code=status_code): + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Blocked by Trend AI Guard. Security violation: {reason}", + should_wrap_with_default_message=False, + status_code=status_code, + blocked_content=True, + ) + case TrendAIProviderFailure(reason=reason, status_code=status_code): + if self.fallback_on_error == "allow": + verbose_proxy_logger.warning( + "Trend AI Guard: %s (status=%s); allowing traffic (fallback_on_error=allow)", + reason, + status_code, + ) + return None + verbose_proxy_logger.error( + "Trend AI Guard: %s (status=%s); blocking traffic (fallback_on_error=block)", reason, status_code + ) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Security Guard Error: {reason}", + should_wrap_with_default_message=False, + status_code=PROVIDER_UNAVAILABLE_STATUS, + ) + case _: + assert_never(result) + + def _record_scan( + self, + result: TrendAIScanResult, + request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict + started_at: float, + *, + event_type: GuardrailEventHooks | None, + redacted: bool | None = None, + ) -> None: + match result: + case TrendAIAllow(redacted_content=redacted_content, masked_entity_count=masked_entity_count): + self._record( + _allow_record(redacted_content is not None if redacted is None else redacted), + "success", + request_data, + started_at, + event_type=event_type, + masked_entity_count=_merge_entity_counts(_NO_ENTITIES, masked_entity_count), + ) + case TrendAIBlock(reason=reason): + self._record( + MappingProxyType({"action": "block", "reason": reason}), + "guardrail_intervened", + request_data, + started_at, + event_type=event_type, + ) + case TrendAIProviderFailure(reason=reason, status_code=status_code): + self._record( + MappingProxyType( + {"description": reason, "status_code": status_code, "fallback_on_error": self.fallback_on_error} + ), + "guardrail_failed_to_respond", + request_data, + started_at, + event_type=event_type, + ) + case _: + assert_never(result) + + def _record_failure( + self, + description: str, + request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict + started_at: float, + *, + event_type: GuardrailEventHooks | None, + ) -> None: + self._record( + MappingProxyType({"description": description}), + "guardrail_failed_to_respond", + request_data, + started_at, + event_type=event_type, + ) + + def _record( + self, + guardrail_json_response: Mapping[str, object], + guardrail_status: GuardrailStatus, + request_data: dict[str, object], # mutable-ok: CustomGuardrail hook contract is a plain dict + started_at: float, + *, + event_type: GuardrailEventHooks | None, + masked_entity_count: Mapping[str, int] = _NO_ENTITIES, + ) -> None: + ended_at: Final = time.time() + self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # base signature takes an untyped dict + guardrail_json_response=dict(guardrail_json_response), # mutable-ok: base signature takes a plain dict + request_data=request_data, + guardrail_status=guardrail_status, + start_time=started_at, + end_time=ended_at, + duration=ended_at - started_at, + event_type=event_type, + masked_entity_count=dict(masked_entity_count) or None, # mutable-ok: base signature takes a plain dict + ) + + def _build_request_headers(self, request_type: str | None = None) -> Mapping[str, str]: + headers: Final = ( + ("TMV1-Application-Name", self.app_name), + ("Authorization", f"Bearer {self.api_key}"), + ("Content-Type", "application/json"), + ("TMV1-Client-Name", TMV1_CLIENT_NAME), + ("TMV1-Client-Version", litellm_version), + ("TMV1-Plugin-Version", PLUGIN_VERSION), + *_optional_header("prefer", "redact-pii,return=representation" if self.mask_pii else None), + *_optional_header("TMV1-Request-Type", request_type), + ) + return MappingProxyType(dict(headers)) + + async def _scan_payload( + self, + payload: Mapping[str, object], + request_type: str | None = None, + ) -> TrendAIScanResult: + try: + response: Final = await self.async_handler.post( + self.api_url, + json=dict(payload), # mutable-ok: httpx takes a plain dict + headers=dict(self._build_request_headers(request_type)), # mutable-ok: httpx takes a plain dict + timeout=self.timeout, + ) + response.raise_for_status() + parsed: Final = TrendAIResponse.model_validate_json(response.content) + except httpx.HTTPStatusError as error: + return TrendAIProviderFailure( + reason="Trend AI Guard returned an HTTP error", + status_code=error.response.status_code, + ) + except (httpx.RequestError, Timeout, asyncio.TimeoutError) as error: + return TrendAIProviderFailure(reason=f"Trend AI Guard request failed: {type(error).__name__}") + except ValidationError: + return TrendAIProviderFailure(reason="Trend AI Guard returned an invalid response") + + normalized_action: Final = parsed.action.strip().lower() + if normalized_action == "block": + reason: Final = ", ".join(parsed.reasons) or parsed.reason or "Content policy violation" + return TrendAIBlock(reason=reason) + if normalized_action != "allow": + return TrendAIProviderFailure(reason=f"Trend AI Guard returned an unsupported action: {parsed.action!r}") + + redacted_content: Final = _extract_redacted_content(request_type, parsed) + if parsed.redacted_request is not None and redacted_content is None: + return TrendAIProviderFailure(reason="Trend AI Guard returned an invalid redacted payload") + + entity_counts: Final = ( + tuple( + sorted( + ( + rule_id, + sum(1 for candidate in parsed.sensitive_information.rules if candidate.id.strip() == rule_id), + ) + for rule_id in frozenset( + rule.id.strip() for rule in parsed.sensitive_information.rules if rule.id.strip() + ) + ) + ) + if parsed.sensitive_information is not None and redacted_content is not None + else () + ) + return TrendAIAllow(redacted_content=redacted_content, masked_entity_count=entity_counts) + + +def _chat_completion_payload(content: str, model: str) -> Mapping[str, object]: + return MappingProxyType(TrendAIChatCompletionPayload.for_content(content, model).model_dump()) + + +def _optional_header(name: str, value: str | None) -> tuple[tuple[str, str], ...]: + return () if value is None else ((name, value),) + + +def _allow_record(redacted: bool) -> Mapping[str, object]: + return MappingProxyType({"action": "allow", "redacted": redacted}) + + +def _is_redacting(result: TrendAIScanResult) -> bool: + return isinstance(result, TrendAIAllow) and result.redacted_content is not None + + +def _window_verdict(scans: Sequence[TrendAIWindowScan]) -> TrendAIScanResult: + """Collapse window scans into one verdict: the first non-allow result, else an allow carrying every entity count.""" + terminal: Final = next((scan.result for scan in scans if not isinstance(scan.result, TrendAIAllow)), None) + if terminal is not None: + return terminal + counts: Final = tuple( + pair for scan in scans if isinstance(scan.result, TrendAIAllow) for pair in scan.result.masked_entity_count + ) + return TrendAIAllow(masked_entity_count=tuple(sorted(_merge_entity_counts(_NO_ENTITIES, counts).items()))) + + +def _merge_entity_counts(counts: Mapping[str, int], additions: Sequence[tuple[str, int]]) -> Mapping[str, int]: + """Per-entity maxima: overlapping scans may report the same entity more than once.""" + entities: Final = frozenset(counts) | frozenset(entity for entity, _ in additions) + return MappingProxyType( + { + entity: max((counts.get(entity, 0), *(count for candidate, count in additions if candidate == entity))) + for entity in entities + } + ) + + +def _merge_window_redactions(content: str, scans: Sequence[TrendAIWindowScan]) -> str | None: + """Fold every window's redaction into ``content``; None if any window cannot be merged positionally.""" + merged = content # rebind-ok: folded across the windows + for scan in scans: + if not isinstance(scan.result, TrendAIAllow) or scan.result.redacted_content is None: + continue + merged = apply_window_redaction(merged, scan.window, scan.result.redacted_content) + if merged is None: + return None + return merged + + +def _build_apply_guardrails_url(api_base: str) -> str: + parsed: Final = urlsplit(api_base.strip()) + normalized_path: Final = parsed.path.rstrip("/") + path_with_product: Final = ( + f"{normalized_path}{_TREND_AI_SECURITY_PATH}" + if parsed.hostname is not None + and parsed.hostname.lower().endswith(".trendmicro.com") + and "/aiSecurity" not in normalized_path + else normalized_path + ) + final_path: Final = ( + path_with_product + if path_with_product.endswith(_APPLY_GUARDRAILS_PATH) + else f"{path_with_product}{_APPLY_GUARDRAILS_PATH}" + ) + return urlunsplit(SplitResult(parsed.scheme, parsed.netloc, final_path, parsed.query, parsed.fragment)) + + +def _extract_redacted_content(request_type: str | None, response: TrendAIResponse) -> str | None: + payload: Final = response.redacted_request + if payload is None: + return None + try: + if request_type == OPENAI_CHAT_COMPLETION_RESPONSE_V1: + parsed_response: Final = TrendAIRedactedResponse.model_validate(payload) + content: Final = " ".join( + choice.message.content + for choice in parsed_response.choices + if choice.message is not None and choice.message.content + ) + return content or None + return TrendAIRedactedPrompt.model_validate(payload).prompt or None + except ValidationError: + return None 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 new file mode 100644 index 00000000000..65b93c53d68 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py @@ -0,0 +1,563 @@ +import json +from collections.abc import Callable, Mapping, Sequence +from typing import Literal + +import httpx +import pytest + +from litellm.exceptions import GuardrailRaisedException +from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler +from litellm.proxy.guardrails.guardrail_hooks.trendai import TrendAIGuardrail +from litellm.proxy.guardrails.guardrail_hooks.trendai._models import ( + TrendAIAllow, + TrendAIBlock, + TrendAIProviderFailure, + TrendAIScanResult, +) +from litellm.proxy.guardrails.guardrail_hooks.trendai._text import utf8_windows +from litellm.proxy.guardrails.guardrail_hooks.trendai.trendai import ( + OPENAI_CHAT_COMPLETION_RESPONSE_V1, + PLUGIN_VERSION, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import ( + CallTypes, + Choices, + GenericGuardrailAPIInputs, + Message, + ModelResponse, +) + + +def _guardrail( + *, + app_name: str | None = None, + fallback_on_error: Literal["block", "allow"] = "block", + mask_pii: bool = True, + 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", + async_handler: httpx.AsyncClient | None = None, + api_base: str = "https://guard.example.com/v3.0/aiSecurity", + event_hook: GuardrailEventHooks = GuardrailEventHooks.pre_call, +) -> TrendAIGuardrail: + return TrendAIGuardrail( + api_key="test-key", + api_base=api_base, + app_name=app_name, + fallback_on_error=fallback_on_error, + mask_pii=mask_pii, + 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, + async_handler=async_handler, + guardrail_name="trendai", + event_hook=event_hook, + ) + + +async def _scan( + guardrail: TrendAIGuardrail, + payload: Mapping[str, object], + request_type: str | None = None, +) -> TrendAIScanResult: + return await guardrail._scan_payload( # pyright: ignore[reportPrivateUsage] # verifies the section 2 transport seam + payload, + request_type=request_type, + ) + + +@pytest.mark.parametrize( + ("api_base", "expected"), + [ + ( + "https://api.xdr.trendmicro.com", + "https://api.xdr.trendmicro.com/v3.0/aiSecurity/applyGuardrails", + ), + ( + "https://api.xdr.trendmicro.com/v3.0/aiSecurity/", + "https://api.xdr.trendmicro.com/v3.0/aiSecurity/applyGuardrails", + ), + ( + "https://self-hosted.example.com/guard/", + "https://self-hosted.example.com/guard/applyGuardrails", + ), + ( + "https://self-hosted.example.com/guard/applyGuardrails", + "https://self-hosted.example.com/guard/applyGuardrails", + ), + ], +) +def test_build_apply_guardrails_url(api_base: str, expected: str) -> None: + assert _guardrail(api_base=api_base).api_url == expected + + +def test_environment_fallbacks(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("TMV1_API_KEY", "env-key") + monkeypatch.setenv("TRENDAI_AI_GUARD_BASE_URL", "https://guard.example.com") + monkeypatch.setenv("TMV1_APPLICATION_NAME", "env-app") + + guardrail = TrendAIGuardrail( + guardrail_name="trendai", + event_hook=GuardrailEventHooks.pre_call, + ) + + assert guardrail.api_key == "env-key" + assert guardrail.api_url == "https://guard.example.com/applyGuardrails" + assert guardrail.app_name == "env-app" + + +def test_explicit_configuration_takes_precedence(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("TMV1_API_KEY", "env-key") + monkeypatch.setenv("TRENDAI_AI_GUARD_BASE_URL", "https://env.example.com") + monkeypatch.setenv("TMV1_APPLICATION_NAME", "env-app") + + guardrail = _guardrail(app_name="configured-app") + + assert guardrail.api_key == "test-key" + assert guardrail.api_url == "https://guard.example.com/v3.0/aiSecurity/applyGuardrails" + assert guardrail.app_name == "configured-app" + + +def test_invalid_stream_configuration_is_rejected() -> None: + with pytest.raises(ValueError, match="stream_overlap_size"): + _guardrail(stream_batch_size=10, stream_overlap_size=10) + + +def test_invalid_timeout_is_rejected() -> None: + with pytest.raises(ValueError, match="timeout"): + _guardrail(timeout=0) + + +@pytest.mark.asyncio +async def test_scan_uses_injected_client_and_required_headers() -> None: + captured_request: httpx.Request | None = None + + async def respond(request: httpx.Request) -> httpx.Response: + nonlocal captured_request + captured_request = request + return httpx.Response(200, json={"action": "allow"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result = await _scan( + _guardrail(async_handler=client, app_name="test-app"), + {"prompt": "hello"}, + request_type=OPENAI_CHAT_COMPLETION_RESPONSE_V1, + ) + + assert isinstance(result, TrendAIAllow) + assert captured_request is not None + assert str(captured_request.url) == "https://guard.example.com/v3.0/aiSecurity/applyGuardrails" + assert captured_request.headers["TMV1-Application-Name"] == "test-app" + assert captured_request.headers["Authorization"] == "Bearer test-key" + assert captured_request.headers["Content-Type"] == "application/json" + assert captured_request.headers["TMV1-Client-Name"] == "litellm" + assert captured_request.headers["TMV1-Plugin-Version"] == PLUGIN_VERSION + assert captured_request.headers["TMV1-Request-Type"] == OPENAI_CHAT_COMPLETION_RESPONSE_V1 + assert captured_request.headers["prefer"] == "redact-pii,return=representation" + assert json.loads(captured_request.content) == {"prompt": "hello"} + + +@pytest.mark.asyncio +async def test_scan_parses_response_redaction_and_entities() -> None: + async def respond(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={ + "action": "allow", + "redactedRequest": {"choices": [{"message": {"content": "email: *****"}}]}, + "sensitiveInformation": {"rules": [{"id": "EMAIL"}, {"id": "EMAIL"}]}, + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result = await _scan( + _guardrail(async_handler=client), + {"choices": []}, + request_type=OPENAI_CHAT_COMPLETION_RESPONSE_V1, + ) + + assert result == TrendAIAllow(redacted_content="email: *****", masked_entity_count=(("EMAIL", 2),)) + + +@pytest.mark.asyncio +async def test_scan_models_block_and_provider_failure_separately() -> None: + responses = iter( + ( + httpx.Response(200, json={"action": "block", "reasons": ["malware", "credential theft"]}), + httpx.Response(200, json={"unexpected": True}), + httpx.Response(503, text="unavailable"), + ) + ) + + async def respond(request: httpx.Request) -> httpx.Response: + return next(responses) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + guardrail = _guardrail(async_handler=client) + blocked = await _scan(guardrail, {"prompt": "blocked"}) + malformed = await _scan(guardrail, {"prompt": "malformed"}) + unavailable = await _scan(guardrail, {"prompt": "unavailable"}) + + assert blocked == TrendAIBlock(reason="malware, credential theft") + assert isinstance(malformed, TrendAIProviderFailure) + assert isinstance(unavailable, TrendAIProviderFailure) + assert unavailable.status_code == 503 + + +def test_guardrail_is_discovered_by_global_registries() -> None: + from litellm.proxy.guardrails.guardrail_registry import ( + guardrail_class_registry, + guardrail_initializer_registry, + ) + + assert guardrail_class_registry["trendai"] is TrendAIGuardrail + assert "trendai" in guardrail_initializer_registry + + +Responder = Callable[[httpx.Request], httpx.Response] + + +def _engine( + *, + block_on: str | None = None, + redact: Mapping[str, str] | None = None, + 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: + body = json.loads(request.content) + text = body["prompt"] if "prompt" in body else body["choices"][0]["message"]["content"] + scanned.append(text) + if fail_on is not None and fail_on in text: + return httpx.Response(503, text="down") + if block_on is not None and block_on in text: + return httpx.Response(200, json={"action": "block", "reasons": ["policy"]}) + 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) + payload = {"prompt": redacted} if "prompt" in body else {"choices": [{"message": {"content": redacted}}]} + return httpx.Response( + 200, + json={ + "action": "allow", + "redactedRequest": payload, + "sensitiveInformation": {"rules": [{"id": entity}]}, + }, + ) + + return respond, scanned + + +def _guardrail_records(request_data: Mapping[str, object]) -> list[dict]: + metadata = request_data["metadata"] + assert isinstance(metadata, dict) + return list(metadata.get("standard_logging_guardrail_information") or []) + + +async def _apply( + guardrail: TrendAIGuardrail, + inputs: GenericGuardrailAPIInputs, + input_type: Literal["request", "response"], +) -> tuple[GenericGuardrailAPIInputs, dict]: + request_data: dict = {"metadata": {}} + result = await guardrail.apply_guardrail(inputs=inputs, request_data=request_data, input_type=input_type) + return result, request_data + + +@pytest.mark.asyncio +async def test_request_scans_only_the_last_user_turn_and_writes_redaction_back() -> None: + respond, scanned = _engine(redact={"a@b.com": "[EMAIL]"}, entity="EMAIL") + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result, request_data = await _apply( + _guardrail(async_handler=client), + { + "texts": ["system prompt", "first turn a@b.com", "reply", "mail a@b.com now"], + "structured_messages": [ + {"role": "system", "content": "system prompt"}, + {"role": "user", "content": "first turn a@b.com"}, + {"role": "assistant", "content": "reply"}, + {"role": "user", "content": "mail a@b.com now"}, + ], + }, + "request", + ) + + assert scanned == ["mail a@b.com now"] + assert result["texts"] == ["system prompt", "first turn a@b.com", "reply", "mail [EMAIL] now"] + (record,) = _guardrail_records(request_data) + assert record["guardrail_status"] == "success" + assert record["masked_entity_count"] == {"EMAIL": 1} + assert record["guardrail_response"] == {"action": "allow", "redacted": True} + + +@pytest.mark.asyncio +async def test_request_redaction_is_split_back_across_multipart_user_content() -> None: + respond, scanned = _engine(redact={"4111": "####"}) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result, _ = await _apply( + _guardrail(async_handler=client), + { + "texts": ["card ", "4111 ok"], + "structured_messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "card "}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,AA=="}}, + {"type": "text", "text": "4111 ok"}, + ], + } + ], + }, + "request", + ) + + assert scanned == ["card 4111 ok"] + assert result["texts"] == ["card ", "#### ok"] + + +@pytest.mark.asyncio +async def test_request_without_user_text_is_not_scanned_or_recorded() -> None: + respond, scanned = _engine() + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + inputs: GenericGuardrailAPIInputs = { + "texts": ["system prompt"], + "structured_messages": [{"role": "system", "content": "system prompt"}], + } + result, request_data = await _apply(_guardrail(async_handler=client), inputs, "request") + + assert scanned == [] + assert result is inputs + assert _guardrail_records(request_data) == [] + + +@pytest.mark.asyncio +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": {}} + with pytest.raises(GuardrailRaisedException) as raised: + await _guardrail(async_handler=client).apply_guardrail( + inputs={"texts": ["build a bomb"]}, + request_data=request_data, + input_type="request", + ) + + assert raised.value.status_code == 400 + assert raised.value.blocked_content is True + assert "policy" in raised.value.message + (record,) = _guardrail_records(request_data) + assert record["guardrail_status"] == "guardrail_intervened" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("fallback_on_error", "raises"), + [("block", True), ("allow", False)], +) +async def test_provider_failure_follows_fallback_policy_and_is_never_an_intervention( + fallback_on_error: Literal["block", "allow"], raises: bool +) -> None: + 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": {}} + inputs: GenericGuardrailAPIInputs = {"texts": ["anything"]} + if raises: + with pytest.raises(GuardrailRaisedException) as raised: + await guardrail.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + assert raised.value.status_code == 503 + assert raised.value.blocked_content is False + else: + result = await guardrail.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request") + assert result is inputs + + (record,) = _guardrail_records(request_data) + assert record["guardrail_status"] == "guardrail_failed_to_respond" + assert record["guardrail_response"]["status_code"] == 503 + assert record["guardrail_response"]["fallback_on_error"] == fallback_on_error + + +@pytest.mark.asyncio +async def test_response_redaction_merges_across_overlapping_windows() -> None: + respond, scanned = _engine(redact={"SECRET": "******", "private": "*******"}) + prefix = "a" * 14 + content = f"{prefix}SECRETprivate" + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result, request_data = await _apply( + _guardrail(async_handler=client, response_content_chunk_size_bytes=20, stream_overlap_size=6), + {"texts": [content], "model": "gpt-5.4"}, + "response", + ) + + assert scanned == [f"{prefix}SECRET", "SECRETprivate"] + assert result["texts"] == [f"{prefix}*************"] + (record,) = _guardrail_records(request_data) + assert record["guardrail_status"] == "success" + assert record["masked_entity_count"] == {"PII": 1} + + +def _redacted_choice(content: str) -> httpx.Response: + return httpx.Response( + 200, json={"action": "allow", "redactedRequest": {"choices": [{"message": {"content": content}}]}} + ) + + +@pytest.mark.asyncio +async def test_later_overlap_scan_cannot_undo_an_earlier_redaction() -> None: + """Window two sees ``SECRET`` in its overlap unmasked, flags only ``tail``, and must not restore ``SECRET``.""" + responses = iter((_redacted_choice("aaaaa******"), _redacted_choice("SECRET####"))) + + async with httpx.AsyncClient(transport=httpx.MockTransport(lambda request: next(responses))) as client: + result, _ = await _apply( + _guardrail(async_handler=client, response_content_chunk_size_bytes=11, stream_overlap_size=6), + {"texts": ["aaaaaSECRETtail"]}, + "response", + ) + + assert result["texts"] == ["aaaaa******####"] + + +@pytest.mark.asyncio +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": {}} + with pytest.raises(GuardrailRaisedException, match="policy"): + await _guardrail( + async_handler=client, response_content_chunk_size_bytes=5, stream_overlap_size=0 + ).apply_guardrail( + inputs={"texts": ["aaaaabbbbbzzzzzccccc"]}, + request_data=request_data, + input_type="response", + ) + + assert scanned == ["aaaaa", "bbbbb", "zzzzz"] + assert _guardrail_records(request_data)[0]["guardrail_status"] == "guardrail_intervened" + + +@pytest.mark.asyncio +async def test_response_scans_every_choice() -> None: + respond, scanned = _engine(redact={"SECRET": "******"}) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result, _ = await _apply( + _guardrail(async_handler=client), + {"texts": ["first SECRET", "second SECRET"]}, + "response", + ) + + assert scanned == ["first SECRET", "second SECRET"] + assert result["texts"] == ["first ******", "second ******"] + + +@pytest.mark.asyncio +async def test_unmergeable_response_redaction_blocks_instead_of_leaking() -> None: + def respond(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"action": "allow", "redactedRequest": {"choices": [{"message": {"content": "short"}}]}}, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + request_data: dict = {"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"]}, + request_data=request_data, + input_type="response", + ) + + assert _guardrail_records(request_data)[-1]["guardrail_status"] == "guardrail_failed_to_respond" + + +@pytest.mark.asyncio +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 = { + "model": "gpt-5.4", + "messages": [ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": [{"type": "text", "text": "mail a@b.com"}]}, + ], + "metadata": {}, + } + out = await OpenAIChatCompletionsHandler().process_input_messages( + data=data, guardrail_to_apply=_guardrail(async_handler=client) + ) + + assert out["messages"][0]["content"] == "be terse" + assert out["messages"][1]["content"][0]["text"] == "mail [EMAIL]" + + +@pytest.mark.parametrize( + ("content", "chunk_size_bytes", "overlap", "expected"), + [ + ("abcdéX", 5, 0, ["abcd", "éX"]), + ("abcéX", 5, 0, ["abcé", "X"]), + ("abcdéXYZ", 5, 2, ["abcd", "cdéX", "éXYZ"]), + ("abcdef", 5, 100, ["abcde", "bcdef"]), + ("", 5, 2, []), + ("é", 1, 0, ["é"]), + ], +) +def test_utf8_windows_respect_byte_limits_and_character_overlap( + content: str, chunk_size_bytes: int, overlap: int, expected: Sequence[str] +) -> None: + windows = utf8_windows(content, chunk_size_bytes=chunk_size_bytes, overlap_chars=overlap) + + assert [window.text for window in windows] == list(expected) + 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]: + response = ModelResponse(choices=[Choices(message=Message(role="assistant", content=assistant_text))]) + kwargs: dict = { + "model": "gpt-5.4", + "messages": [{"role": "user", "content": user_text}], + "litellm_call_id": "call-1", + "litellm_params": {"metadata": {}}, + "optional_params": {}, + "standard_logging_object": {"guardrail_information": None}, + } + return kwargs, response + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("scope", "expected_scans"), + [ + ("request", ["user a@b.com"]), + ("response", ["assistant SECRET"]), + ("both", ["user a@b.com", "assistant SECRET"]), + ], +) +async def test_logging_only_scan_scope_selects_which_side_is_scanned( + scope: Literal["request", "response", "both"], expected_scans: Sequence[str] +) -> None: + respond, scanned = _engine(redact={"a@b.com": "[EMAIL]", "SECRET": "******"}) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + guardrail = _guardrail( + async_handler=client, logging_only_scan=scope, event_hook=GuardrailEventHooks.logging_only + ) + kwargs, response = _logged_call("user a@b.com", "assistant SECRET") + out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) + + assert scanned == list(expected_scans) + assert out_kwargs["messages"] == [{"role": "user", "content": "user a@b.com"}] + assert out_response is response + assert response.choices[0].message.content == "assistant SECRET" + entries = out_kwargs["standard_logging_object"]["guardrail_information"] + assert [entry["guardrail_status"] for entry in entries] == ["success"] * len(expected_scans) + assert all(entry["guardrail_mode"] == "logging_only" for entry in entries) From fc65a9a3334b91c432e72824a2b4db3716e8962d Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Tue, 15 Sep 2026 13:13:08 -0400 Subject: [PATCH 2/7] fix(guardrails): remove TrendAI mask setting --- .../guardrails/guardrail_hooks/trendai/__init__.py | 1 - .../guardrails/guardrail_hooks/trendai/_models.py | 1 - .../guardrails/guardrail_hooks/trendai/trendai.py | 8 +++----- .../guardrail_hooks/trendai/test_trendai.py | 11 +++++++++-- 4 files changed, 12 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py index b6ca0d3a7c7..7803ac0baea 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py @@ -28,7 +28,6 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" api_base=litellm_params.api_base, app_name=settings.app_name, fallback_on_error=settings.fallback_on_error, - mask_pii=settings.mask_pii, timeout=settings.timeout, stream_batch_size=settings.stream_batch_size, stream_overlap_size=settings.stream_overlap_size, diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py index 2efcbb52c75..0acaf1fb0d3 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py @@ -11,7 +11,6 @@ from pydantic import BaseModel, ConfigDict, Field class TrendAISettings(BaseModel): app_name: str | None = None fallback_on_error: Literal["block", "allow"] = "block" - mask_pii: bool = True timeout: float = 5.0 stream_batch_size: int = 2048 stream_overlap_size: int = 256 diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py index 4ed0e54abe8..c6753738ffc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -74,7 +74,6 @@ class TrendAIGuardrail(CustomGuardrail): api_base: str | None = None, app_name: str | None = None, fallback_on_error: Literal["block", "allow"] = "block", - mask_pii: bool = True, timeout: float = 5.0, stream_batch_size: int = 2048, stream_overlap_size: int = 256, @@ -114,7 +113,6 @@ class TrendAIGuardrail(CustomGuardrail): self.api_url: str = _build_apply_guardrails_url(resolved_api_base) 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.mask_pii: bool = mask_pii self.timeout: float = timeout self.stream_batch_size: int = stream_batch_size self.stream_overlap_size: int = stream_overlap_size @@ -129,8 +127,8 @@ class TrendAIGuardrail(CustomGuardrail): event_hook=event_hook, default_on=default_on, supported_event_hooks=self.get_supported_event_hooks(), - mask_request_content=mask_pii, - mask_response_content=mask_pii, + mask_request_content=True, + mask_response_content=True, ) @classmethod @@ -369,7 +367,7 @@ class TrendAIGuardrail(CustomGuardrail): ("TMV1-Client-Name", TMV1_CLIENT_NAME), ("TMV1-Client-Version", litellm_version), ("TMV1-Plugin-Version", PLUGIN_VERSION), - *_optional_header("prefer", "redact-pii,return=representation" if self.mask_pii else None), + ("prefer", "redact-pii,return=representation"), *_optional_header("TMV1-Request-Type", request_type), ) return MappingProxyType(dict(headers)) 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 65b93c53d68..4ae752cecd4 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 @@ -33,7 +33,6 @@ def _guardrail( *, app_name: str | None = None, fallback_on_error: Literal["block", "allow"] = "block", - mask_pii: bool = True, timeout: float = 5.0, stream_batch_size: int = 2048, stream_overlap_size: int = 256, @@ -48,7 +47,6 @@ def _guardrail( api_base=api_base, app_name=app_name, fallback_on_error=fallback_on_error, - mask_pii=mask_pii, timeout=timeout, stream_batch_size=stream_batch_size, stream_overlap_size=stream_overlap_size, @@ -133,6 +131,15 @@ def test_invalid_timeout_is_rejected() -> None: _guardrail(timeout=0) +def test_pii_masking_is_not_configurable_on_the_guardrail() -> None: + with pytest.raises(TypeError, match="mask_pii"): + TrendAIGuardrail( + api_key="test-key", + api_base="https://guard.example.com/v3.0/aiSecurity", + mask_pii=False, + ) + + @pytest.mark.asyncio async def test_scan_uses_injected_client_and_required_headers() -> None: captured_request: httpx.Request | None = None From d011f18d65abc5e5fb4050a16103de19cf6063ef Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 15:19:58 -0400 Subject: [PATCH 3/7] 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", From b6e4e45a39d4360529257f083cc791ada8cfc1e4 Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 15:29:02 -0400 Subject: [PATCH 4/7] fix(guardrails): remove unused TrendAI logging scan scope --- .../guardrail_hooks/trendai/__init__.py | 1 - .../guardrail_hooks/trendai/_models.py | 1 - .../guardrail_hooks/trendai/trendai.py | 7 ------ .../guardrail_hooks/trendai/test_trendai.py | 22 ++++--------------- 4 files changed, 4 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py index 27b2d1206f7..a506dab07bb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py @@ -31,7 +31,6 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" timeout=settings.timeout, 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, guardrail_name=guardrail_name, event_hook=_normalize_event_hook(litellm_params.mode), default_on=litellm_params.default_on is True, diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py index b3ece6e2973..052d8e50630 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py @@ -15,7 +15,6 @@ class TrendAISettings(BaseModel): timeout: float = 5.0 stream_overlap_size: int = 256 response_content_chunk_size_bytes: int = 49_500 - logging_only_scan: Literal["request", "response", "both"] = "both" class TrendAIChatMessage(BaseModel): diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py index c55dfe1df58..d6ceea50f48 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -80,7 +80,6 @@ class TrendAIGuardrail(CustomGuardrail): timeout: float = 5.0, stream_overlap_size: int = 256, response_content_chunk_size_bytes: int = RESPONSE_CONTENT_CHUNK_SIZE_BYTES, - logging_only_scan: Literal["request", "response", "both"] = "both", async_handler: _AsyncHTTPClient | None = None, guardrail_name: str | None = None, event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, @@ -100,8 +99,6 @@ class TrendAIGuardrail(CustomGuardrail): ) if fallback_on_error not in ("block", "allow"): raise ValueError("fallback_on_error must be 'block' or 'allow'") - if logging_only_scan not in ("request", "response", "both"): - raise ValueError("logging_only_scan must be 'request', 'response', or 'both'") if timeout <= 0: raise ValueError("timeout must be greater than zero") if stream_overlap_size < 0: @@ -116,7 +113,6 @@ class TrendAIGuardrail(CustomGuardrail): self.timeout: float = timeout 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 self.async_handler: _AsyncHTTPClient = async_handler or get_async_httpx_client( llm_provider=httpxSpecialProvider.GuardrailCallback ) @@ -137,9 +133,6 @@ class TrendAIGuardrail(CustomGuardrail): GuardrailEventHooks.logging_only, ] - def logging_only_scan_scope(self) -> Literal["request", "response", "both"]: - return self.logging_only_scan - @log_guardrail_information async def apply_guardrail( self, 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 a2b197da315..250bc30adbd 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 @@ -37,7 +37,6 @@ def _guardrail( timeout: float = 5.0, stream_overlap_size: int = 256, response_content_chunk_size_bytes: int = 49_500, - logging_only_scan: Literal["request", "response", "both"] = "both", async_handler: httpx.AsyncClient | None = None, api_base: str = "https://guard.example.com/v3.0/aiSecurity", event_hook: GuardrailEventHooks = GuardrailEventHooks.pre_call, @@ -50,7 +49,6 @@ def _guardrail( timeout=timeout, stream_overlap_size=stream_overlap_size, response_content_chunk_size_bytes=response_content_chunk_size_bytes, - logging_only_scan=logging_only_scan, async_handler=async_handler, guardrail_name="trendai", event_hook=event_hook, @@ -538,29 +536,17 @@ def _logged_call(user_text: str, assistant_text: str) -> tuple[dict[str, object] @pytest.mark.asyncio -@pytest.mark.parametrize( - ("scope", "expected_scans"), - [ - ("request", ["user a@b.com"]), - ("response", ["assistant SECRET"]), - ("both", ["user a@b.com", "assistant SECRET"]), - ], -) -async def test_logging_only_scan_scope_selects_which_side_is_scanned( - scope: Literal["request", "response", "both"], expected_scans: Sequence[str] -) -> None: +async def test_logging_only_scans_both_sides_without_modifying_them() -> None: respond, scanned = _engine(redact={"a@b.com": "[EMAIL]", "SECRET": "******"}) async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: - guardrail = _guardrail( - async_handler=client, logging_only_scan=scope, event_hook=GuardrailEventHooks.logging_only - ) + guardrail = _guardrail(async_handler=client, event_hook=GuardrailEventHooks.logging_only) kwargs, response = _logged_call("user a@b.com", "assistant SECRET") out_kwargs, out_response = await guardrail.async_logging_hook(kwargs, response, CallTypes.acompletion.value) - assert scanned == list(expected_scans) + assert scanned == ["user a@b.com", "assistant SECRET"] assert out_kwargs["messages"] == [{"role": "user", "content": "user a@b.com"}] assert out_response is response assert response.choices[0].message.content == "assistant SECRET" entries = out_kwargs["standard_logging_object"]["guardrail_information"] - assert [entry["guardrail_status"] for entry in entries] == ["success"] * len(expected_scans) + assert [entry["guardrail_status"] for entry in entries] == ["success", "success"] assert all(entry["guardrail_mode"] == "logging_only" for entry in entries) From 721d1efe03882a3ccc2e0f3c75773ecee28ed559 Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 16:01:49 -0400 Subject: [PATCH 5/7] fix(guardrails): address TrendAI lint budget --- .../guardrails/guardrail_hooks/trendai/__init__.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py index a506dab07bb..518089f0d02 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py @@ -16,7 +16,7 @@ def _normalize_event_hook( if isinstance(mode, str): return GuardrailEventHooks(mode) if isinstance(mode, list): - return [GuardrailEventHooks(item) for item in mode] + return [GuardrailEventHooks(item) for item in mode] # mutable-ok: guardrail event hooks require a list return mode @@ -41,7 +41,11 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" return callback -guardrail_initializer_registry: Final = {GUARDRAIL_NAME: initialize_guardrail} -guardrail_class_registry: Final = {GUARDRAIL_NAME: TrendAIGuardrail} +guardrail_initializer_registry: Final = { # mutable-ok: registry discovery requires a dict + GUARDRAIL_NAME: initialize_guardrail +} +guardrail_class_registry: Final = { # mutable-ok: registry discovery requires a dict + GUARDRAIL_NAME: TrendAIGuardrail +} __all__ = ("TrendAIGuardrail",) From 477a58a91f26bed0c3b2711bd628686b670aaa82 Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 16:47:52 -0400 Subject: [PATCH 6/7] refactor(guardrails): reuse secret and message text helpers for TrendAI --- .../guardrail_hooks/trendai/_text.py | 48 +++---------------- .../guardrail_hooks/trendai/trendai.py | 5 +- .../guardrail_hooks/trendai/test_trendai.py | 41 +++++++++++++++- 3 files changed, 49 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py index 4cf09afabb8..c807dd0fb5b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py @@ -3,38 +3,14 @@ # Licensed under the Apache License, Version 2.0. See LICENSE.txt in this directory. from collections.abc import Iterator, Sequence -from typing import Final, Literal - -from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing import Final +from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts from litellm.types.llms.openai import AllMessageValues from ._models import TrendAIRequestPrompt, TrendAITextWindow -class _TextPart(BaseModel): - model_config = ConfigDict(extra="ignore") - - type: Literal["text"] - text: str - - -class _OtherPart(BaseModel): - model_config = ConfigDict(extra="ignore") - - type: str - - -class _UserMessage(BaseModel): - model_config = ConfigDict(extra="ignore") - - role: str - content: str | tuple[_TextPart | _OtherPart, ...] | None = None - - -_MESSAGES: Final = TypeAdapter(tuple[_UserMessage, ...]) - - def utf8_windows(content: str, *, chunk_size_bytes: int, overlap_chars: int) -> tuple[TrendAITextWindow, ...]: """Split ``content`` into windows of at most ``chunk_size_bytes`` UTF-8 bytes. @@ -85,23 +61,11 @@ def apply_window_redaction(content: str, window: TrendAITextWindow, redacted: st return f"{content[: window.start]}{merged}{content[end:]}" -def _text_parts(content: str | tuple[_TextPart | _OtherPart, ...] | None) -> tuple[str, ...] | None: - if content is None: - return None - if isinstance(content, str): - return (content,) - return tuple(part.text for part in content if isinstance(part, _TextPart)) - - def _last_user_text_parts(structured_messages: Sequence[AllMessageValues]) -> tuple[str, ...] | None: - try: - messages: Final = _MESSAGES.validate_python(structured_messages) - except ValidationError: - return None - last_user_message: Final = next((message for message in reversed(messages) if message.role == "user"), None) - if last_user_message is None: - return None - return _text_parts(last_user_message.content) + last_user_message: Final = next( + (message for message in reversed(structured_messages) if message.get("role") == "user"), None + ) + return message_slot_texts(last_user_message) if last_user_message is not None else None def _last_occurrence(texts: Sequence[str], parts: Sequence[str]) -> int | None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py index d6ceea50f48..2f7f178f837 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -25,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # legacy client factory has an untyped params map httpxSpecialProvider, ) +from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus @@ -85,13 +86,13 @@ class TrendAIGuardrail(CustomGuardrail): event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, ) -> None: - resolved_api_key: Final = api_key or os.environ.get("TMV1_API_KEY") + resolved_api_key: Final = api_key or get_secret_str("TMV1_API_KEY") if not resolved_api_key: raise ValueError( "Trend AI Guard requires an API key. Pass api_key or set the TMV1_API_KEY environment variable." ) - resolved_api_base: Final = api_base or os.environ.get("TRENDAI_AI_GUARD_BASE_URL") + resolved_api_base: Final = api_base or get_secret_str("TRENDAI_AI_GUARD_BASE_URL") if not resolved_api_base: raise ValueError( "Trend AI Guard requires an API base URL. Pass api_base or set the " 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 250bc30adbd..afab0e8ecf8 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,6 +1,6 @@ import json -from functools import reduce from collections.abc import Callable, Mapping, Sequence +from functools import reduce from typing import Literal import httpx @@ -106,6 +106,28 @@ def test_environment_fallbacks(monkeypatch: pytest.MonkeyPatch) -> None: assert guardrail.app_name == "env-app" +def test_secret_manager_configuration_is_used(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm + from litellm.secret_managers import main as secrets + + monkeypatch.delenv("TMV1_API_KEY", raising=False) + monkeypatch.delenv("TRENDAI_AI_GUARD_BASE_URL", raising=False) + monkeypatch.setattr(litellm, "secret_manager_client", object()) + monkeypatch.setattr(litellm, "_key_management_settings", None) + monkeypatch.setattr(litellm, "_key_management_system", None) + monkeypatch.setattr(secrets, "_should_read_secret_from_secret_manager", lambda: True) + managed = { + "TMV1_API_KEY": "managed-key", + "TRENDAI_AI_GUARD_BASE_URL": "https://managed.example.com", + } + monkeypatch.setattr(secrets, "get_secret_from_manager", lambda **kwargs: managed[kwargs["secret_name"]]) + + guardrail = TrendAIGuardrail(guardrail_name="trendai", event_hook=GuardrailEventHooks.pre_call) + + assert guardrail.api_key == "managed-key" + assert guardrail.api_url == "https://managed.example.com/applyGuardrails" + + def test_explicit_configuration_takes_precedence(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TMV1_API_KEY", "env-key") monkeypatch.setenv("TRENDAI_AI_GUARD_BASE_URL", "https://env.example.com") @@ -328,6 +350,23 @@ async def test_request_redaction_is_split_back_across_multipart_user_content() - assert result["texts"] == ["card ", "#### ok"] +@pytest.mark.asyncio +async def test_request_scans_text_slots_from_normalized_message_parts() -> None: + respond, scanned = _engine(redact={"a@b.com": "[EMAIL]"}) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result, _ = await _apply( + _guardrail(async_handler=client), + { + "texts": ["mail a@b.com"], + "structured_messages": [{"role": "user", "content": [{"type": "input_text", "text": "mail a@b.com"}]}], + }, + "request", + ) + + assert scanned == ["mail a@b.com"] + assert result["texts"] == ["mail [EMAIL]"] + + @pytest.mark.asyncio async def test_request_without_user_text_is_not_scanned_or_recorded() -> None: respond, scanned = _engine() From 1d0bcdbec71d76ce42508cc0f6a916dac23d3ac4 Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 16:53:46 -0400 Subject: [PATCH 7/7] fix(guardrails): restore documented TrendAI environment fallback --- .../guardrail_hooks/trendai/trendai.py | 5 ++--- .../guardrail_hooks/trendai/test_trendai.py | 22 ------------------- 2 files changed, 2 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py index 2f7f178f837..d6ceea50f48 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -25,7 +25,6 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # legacy client factory has an untyped params map httpxSpecialProvider, ) -from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus @@ -86,13 +85,13 @@ class TrendAIGuardrail(CustomGuardrail): event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, ) -> None: - resolved_api_key: Final = api_key or get_secret_str("TMV1_API_KEY") + resolved_api_key: Final = api_key or os.environ.get("TMV1_API_KEY") if not resolved_api_key: raise ValueError( "Trend AI Guard requires an API key. Pass api_key or set the TMV1_API_KEY environment variable." ) - resolved_api_base: Final = api_base or get_secret_str("TRENDAI_AI_GUARD_BASE_URL") + resolved_api_base: Final = api_base or os.environ.get("TRENDAI_AI_GUARD_BASE_URL") if not resolved_api_base: raise ValueError( "Trend AI Guard requires an API base URL. Pass api_base or set the " 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 afab0e8ecf8..1a17db0a8bc 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 @@ -106,28 +106,6 @@ def test_environment_fallbacks(monkeypatch: pytest.MonkeyPatch) -> None: assert guardrail.app_name == "env-app" -def test_secret_manager_configuration_is_used(monkeypatch: pytest.MonkeyPatch) -> None: - import litellm - from litellm.secret_managers import main as secrets - - monkeypatch.delenv("TMV1_API_KEY", raising=False) - monkeypatch.delenv("TRENDAI_AI_GUARD_BASE_URL", raising=False) - monkeypatch.setattr(litellm, "secret_manager_client", object()) - monkeypatch.setattr(litellm, "_key_management_settings", None) - monkeypatch.setattr(litellm, "_key_management_system", None) - monkeypatch.setattr(secrets, "_should_read_secret_from_secret_manager", lambda: True) - managed = { - "TMV1_API_KEY": "managed-key", - "TRENDAI_AI_GUARD_BASE_URL": "https://managed.example.com", - } - monkeypatch.setattr(secrets, "get_secret_from_manager", lambda **kwargs: managed[kwargs["secret_name"]]) - - guardrail = TrendAIGuardrail(guardrail_name="trendai", event_hook=GuardrailEventHooks.pre_call) - - assert guardrail.api_key == "managed-key" - assert guardrail.api_url == "https://managed.example.com/applyGuardrails" - - def test_explicit_configuration_takes_precedence(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TMV1_API_KEY", "env-key") monkeypatch.setenv("TRENDAI_AI_GUARD_BASE_URL", "https://env.example.com")