This commit is contained in:
Oliver Fei 2026-09-28 13:00:17 -04:00 • committed by GitHub
commit 592caa3b93
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1576 additions and 0 deletions

View file

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

View file

@ -0,0 +1,51 @@
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] # mutable-ok: guardrail event hooks require a list
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,
timeout=settings.timeout,
stream_overlap_size=settings.stream_overlap_size,
response_content_chunk_size_bytes=settings.response_content_chunk_size_bytes,
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 = { # 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",)

View file

@ -0,0 +1,132 @@
# 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 Mapping
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"
timeout: float = 5.0
stream_overlap_size: int = 256
response_content_chunk_size_bytes: int = 49_500
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: Mapping[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

View file

@ -0,0 +1,121 @@
# 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
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts
from litellm.types.llms.openai import AllMessageValues
from ._models import TrendAIRequestPrompt, TrendAITextWindow
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 _last_user_text_parts(structured_messages: Sequence[AllMessageValues]) -> tuple[str, ...] | None:
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:
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 :])

View file

@ -0,0 +1,502 @@
# 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, 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,
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,
)
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",
timeout: float = 5.0,
stream_overlap_size: int = 256,
response_content_chunk_size_bytes: int = RESPONSE_CONTENT_CHUNK_SIZE_BYTES,
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 timeout <= 0:
raise ValueError("timeout must be greater than zero")
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")
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.timeout: float = timeout
self.stream_overlap_size: int = stream_overlap_size
self.response_content_chunk_size_bytes: int = response_content_chunk_size_bytes
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(),
)
@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,
]
@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: "LiteLLMLoggingObj | None" = 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),
("prefer", "redact-pii,return=representation"),
*_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

View file

@ -0,0 +1,569 @@
import json
from collections.abc import Callable, Mapping, Sequence
from functools import reduce
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",
timeout: float = 5.0,
stream_overlap_size: int = 256,
response_content_chunk_size_bytes: int = 49_500,
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,
timeout=timeout,
stream_overlap_size=stream_overlap_size,
response_content_chunk_size_bytes=response_content_chunk_size_bytes,
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_negative_stream_overlap_is_rejected() -> None:
with pytest.raises(ValueError, match="stream_overlap_size"):
_guardrail(stream_overlap_size=-1)
def test_invalid_timeout_is_rejected() -> None:
with pytest.raises(ValueError, match="timeout"):
_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
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]]:
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 = 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,
json={
"action": "allow",
"redactedRequest": payload,
"sensitiveInformation": {"rules": [{"id": entity}]},
},
)
return respond, scanned
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 [])
async def _apply(
guardrail: TrendAIGuardrail,
inputs: GenericGuardrailAPIInputs,
input_type: Literal["request", "response"],
) -> 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
@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_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()
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[str, object] = {"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[str, object] = {"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[str, object] = {"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[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"]},
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[str, object] = {
"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[str, object], ModelResponse]:
response = ModelResponse(choices=[Choices(message=Message(role="assistant", content=assistant_text))])
kwargs: dict[str, object] = {
"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
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, 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 == ["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", "success"]
assert all(entry["guardrail_mode"] == "logging_only" for entry in entries)