refactor(guardrails): reuse secret and message text helpers for TrendAI

This commit is contained in:
Oliver Fei 2026-09-24 16:47:52 -04:00
parent 721d1efe03
commit 477a58a91f
3 changed files with 49 additions and 45 deletions

View file

@ -3,38 +3,14 @@
# Licensed under the Apache License, Version 2.0. See LICENSE.txt in this directory.
from collections.abc import Iterator, Sequence
from typing import Final, Literal
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
from typing import Final
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts
from litellm.types.llms.openai import AllMessageValues
from ._models import TrendAIRequestPrompt, TrendAITextWindow
class _TextPart(BaseModel):
model_config = ConfigDict(extra="ignore")
type: Literal["text"]
text: str
class _OtherPart(BaseModel):
model_config = ConfigDict(extra="ignore")
type: str
class _UserMessage(BaseModel):
model_config = ConfigDict(extra="ignore")
role: str
content: str | tuple[_TextPart | _OtherPart, ...] | None = None
_MESSAGES: Final = TypeAdapter(tuple[_UserMessage, ...])
def utf8_windows(content: str, *, chunk_size_bytes: int, overlap_chars: int) -> tuple[TrendAITextWindow, ...]:
"""Split ``content`` into windows of at most ``chunk_size_bytes`` UTF-8 bytes.
@ -85,23 +61,11 @@ def apply_window_redaction(content: str, window: TrendAITextWindow, redacted: st
return f"{content[: window.start]}{merged}{content[end:]}"
def _text_parts(content: str | tuple[_TextPart | _OtherPart, ...] | None) -> tuple[str, ...] | None:
if content is None:
return None
if isinstance(content, str):
return (content,)
return tuple(part.text for part in content if isinstance(part, _TextPart))
def _last_user_text_parts(structured_messages: Sequence[AllMessageValues]) -> tuple[str, ...] | None:
try:
messages: Final = _MESSAGES.validate_python(structured_messages)
except ValidationError:
return None
last_user_message: Final = next((message for message in reversed(messages) if message.role == "user"), None)
if last_user_message is None:
return None
return _text_parts(last_user_message.content)
last_user_message: Final = next(
(message for message in reversed(structured_messages) if message.get("role") == "user"), None
)
return message_slot_texts(last_user_message) if last_user_message is not None else None
def _last_occurrence(texts: Sequence[str], parts: Sequence[str]) -> int | None:

View file

@ -25,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # legacy client factory has an untyped params map
httpxSpecialProvider,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus
@ -85,13 +86,13 @@ class TrendAIGuardrail(CustomGuardrail):
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None,
default_on: bool = False,
) -> None:
resolved_api_key: Final = api_key or os.environ.get("TMV1_API_KEY")
resolved_api_key: Final = api_key or get_secret_str("TMV1_API_KEY")
if not resolved_api_key:
raise ValueError(
"Trend AI Guard requires an API key. Pass api_key or set the TMV1_API_KEY environment variable."
)
resolved_api_base: Final = api_base or os.environ.get("TRENDAI_AI_GUARD_BASE_URL")
resolved_api_base: Final = api_base or get_secret_str("TRENDAI_AI_GUARD_BASE_URL")
if not resolved_api_base:
raise ValueError(
"Trend AI Guard requires an API base URL. Pass api_base or set the "

View file

@ -1,6 +1,6 @@
import json
from functools import reduce
from collections.abc import Callable, Mapping, Sequence
from functools import reduce
from typing import Literal
import httpx
@ -106,6 +106,28 @@ def test_environment_fallbacks(monkeypatch: pytest.MonkeyPatch) -> None:
assert guardrail.app_name == "env-app"
def test_secret_manager_configuration_is_used(monkeypatch: pytest.MonkeyPatch) -> None:
import litellm
from litellm.secret_managers import main as secrets
monkeypatch.delenv("TMV1_API_KEY", raising=False)
monkeypatch.delenv("TRENDAI_AI_GUARD_BASE_URL", raising=False)
monkeypatch.setattr(litellm, "secret_manager_client", object())
monkeypatch.setattr(litellm, "_key_management_settings", None)
monkeypatch.setattr(litellm, "_key_management_system", None)
monkeypatch.setattr(secrets, "_should_read_secret_from_secret_manager", lambda: True)
managed = {
"TMV1_API_KEY": "managed-key",
"TRENDAI_AI_GUARD_BASE_URL": "https://managed.example.com",
}
monkeypatch.setattr(secrets, "get_secret_from_manager", lambda **kwargs: managed[kwargs["secret_name"]])
guardrail = TrendAIGuardrail(guardrail_name="trendai", event_hook=GuardrailEventHooks.pre_call)
assert guardrail.api_key == "managed-key"
assert guardrail.api_url == "https://managed.example.com/applyGuardrails"
def test_explicit_configuration_takes_precedence(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("TMV1_API_KEY", "env-key")
monkeypatch.setenv("TRENDAI_AI_GUARD_BASE_URL", "https://env.example.com")
@ -328,6 +350,23 @@ async def test_request_redaction_is_split_back_across_multipart_user_content() -
assert result["texts"] == ["card ", "#### ok"]
@pytest.mark.asyncio
async def test_request_scans_text_slots_from_normalized_message_parts() -> None:
respond, scanned = _engine(redact={"a@b.com": "[EMAIL]"})
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
result, _ = await _apply(
_guardrail(async_handler=client),
{
"texts": ["mail a@b.com"],
"structured_messages": [{"role": "user", "content": [{"type": "input_text", "text": "mail a@b.com"}]}],
},
"request",
)
assert scanned == ["mail a@b.com"]
assert result["texts"] == ["mail [EMAIL]"]
@pytest.mark.asyncio
async def test_request_without_user_text_is_not_scanned_or_recorded() -> None:
respond, scanned = _engine()