mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 1d0bcdbec7 into 6f5ad78a1f
This commit is contained in:
commit
592caa3b93
6 changed files with 1576 additions and 0 deletions
201
litellm/proxy/guardrails/guardrail_hooks/trendai/LICENSE.txt
Normal file
201
litellm/proxy/guardrails/guardrail_hooks/trendai/LICENSE.txt
Normal 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.
|
||||
51
litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py
Normal file
51
litellm/proxy/guardrails/guardrail_hooks/trendai/__init__.py
Normal 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",)
|
||||
132
litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py
Normal file
132
litellm/proxy/guardrails/guardrail_hooks/trendai/_models.py
Normal 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
|
||||
121
litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py
Normal file
121
litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py
Normal 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 :])
|
||||
502
litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py
Normal file
502
litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py
Normal 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
|
||||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue