From 477a58a91f26bed0c3b2711bd628686b670aaa82 Mon Sep 17 00:00:00 2001 From: Oliver Fei Date: Thu, 24 Sep 2026 16:47:52 -0400 Subject: [PATCH] refactor(guardrails): reuse secret and message text helpers for TrendAI --- .../guardrail_hooks/trendai/_text.py | 48 +++---------------- .../guardrail_hooks/trendai/trendai.py | 5 +- .../guardrail_hooks/trendai/test_trendai.py | 41 +++++++++++++++- 3 files changed, 49 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py index 4cf09afabb8..c807dd0fb5b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/_text.py @@ -3,38 +3,14 @@ # Licensed under the Apache License, Version 2.0. See LICENSE.txt in this directory. from collections.abc import Iterator, Sequence -from typing import Final, Literal - -from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing import Final +from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts from litellm.types.llms.openai import AllMessageValues from ._models import TrendAIRequestPrompt, TrendAITextWindow -class _TextPart(BaseModel): - model_config = ConfigDict(extra="ignore") - - type: Literal["text"] - text: str - - -class _OtherPart(BaseModel): - model_config = ConfigDict(extra="ignore") - - type: str - - -class _UserMessage(BaseModel): - model_config = ConfigDict(extra="ignore") - - role: str - content: str | tuple[_TextPart | _OtherPart, ...] | None = None - - -_MESSAGES: Final = TypeAdapter(tuple[_UserMessage, ...]) - - def utf8_windows(content: str, *, chunk_size_bytes: int, overlap_chars: int) -> tuple[TrendAITextWindow, ...]: """Split ``content`` into windows of at most ``chunk_size_bytes`` UTF-8 bytes. @@ -85,23 +61,11 @@ def apply_window_redaction(content: str, window: TrendAITextWindow, redacted: st return f"{content[: window.start]}{merged}{content[end:]}" -def _text_parts(content: str | tuple[_TextPart | _OtherPart, ...] | None) -> tuple[str, ...] | None: - if content is None: - return None - if isinstance(content, str): - return (content,) - return tuple(part.text for part in content if isinstance(part, _TextPart)) - - def _last_user_text_parts(structured_messages: Sequence[AllMessageValues]) -> tuple[str, ...] | None: - try: - messages: Final = _MESSAGES.validate_python(structured_messages) - except ValidationError: - return None - last_user_message: Final = next((message for message in reversed(messages) if message.role == "user"), None) - if last_user_message is None: - return None - return _text_parts(last_user_message.content) + last_user_message: Final = next( + (message for message in reversed(structured_messages) if message.get("role") == "user"), None + ) + return message_slot_texts(last_user_message) if last_user_message is not None else None def _last_occurrence(texts: Sequence[str], parts: Sequence[str]) -> int | None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py index d6ceea50f48..2f7f178f837 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/trendai/trendai.py @@ -25,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # legacy client factory has an untyped params map httpxSpecialProvider, ) +from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus @@ -85,13 +86,13 @@ class TrendAIGuardrail(CustomGuardrail): event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, ) -> None: - resolved_api_key: Final = api_key or os.environ.get("TMV1_API_KEY") + resolved_api_key: Final = api_key or get_secret_str("TMV1_API_KEY") if not resolved_api_key: raise ValueError( "Trend AI Guard requires an API key. Pass api_key or set the TMV1_API_KEY environment variable." ) - resolved_api_base: Final = api_base or os.environ.get("TRENDAI_AI_GUARD_BASE_URL") + resolved_api_base: Final = api_base or get_secret_str("TRENDAI_AI_GUARD_BASE_URL") if not resolved_api_base: raise ValueError( "Trend AI Guard requires an API base URL. Pass api_base or set the " diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py index 250bc30adbd..afab0e8ecf8 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/trendai/test_trendai.py @@ -1,6 +1,6 @@ import json -from functools import reduce from collections.abc import Callable, Mapping, Sequence +from functools import reduce from typing import Literal import httpx @@ -106,6 +106,28 @@ def test_environment_fallbacks(monkeypatch: pytest.MonkeyPatch) -> None: assert guardrail.app_name == "env-app" +def test_secret_manager_configuration_is_used(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm + from litellm.secret_managers import main as secrets + + monkeypatch.delenv("TMV1_API_KEY", raising=False) + monkeypatch.delenv("TRENDAI_AI_GUARD_BASE_URL", raising=False) + monkeypatch.setattr(litellm, "secret_manager_client", object()) + monkeypatch.setattr(litellm, "_key_management_settings", None) + monkeypatch.setattr(litellm, "_key_management_system", None) + monkeypatch.setattr(secrets, "_should_read_secret_from_secret_manager", lambda: True) + managed = { + "TMV1_API_KEY": "managed-key", + "TRENDAI_AI_GUARD_BASE_URL": "https://managed.example.com", + } + monkeypatch.setattr(secrets, "get_secret_from_manager", lambda **kwargs: managed[kwargs["secret_name"]]) + + guardrail = TrendAIGuardrail(guardrail_name="trendai", event_hook=GuardrailEventHooks.pre_call) + + assert guardrail.api_key == "managed-key" + assert guardrail.api_url == "https://managed.example.com/applyGuardrails" + + def test_explicit_configuration_takes_precedence(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TMV1_API_KEY", "env-key") monkeypatch.setenv("TRENDAI_AI_GUARD_BASE_URL", "https://env.example.com") @@ -328,6 +350,23 @@ async def test_request_redaction_is_split_back_across_multipart_user_content() - assert result["texts"] == ["card ", "#### ok"] +@pytest.mark.asyncio +async def test_request_scans_text_slots_from_normalized_message_parts() -> None: + respond, scanned = _engine(redact={"a@b.com": "[EMAIL]"}) + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + result, _ = await _apply( + _guardrail(async_handler=client), + { + "texts": ["mail a@b.com"], + "structured_messages": [{"role": "user", "content": [{"type": "input_text", "text": "mail a@b.com"}]}], + }, + "request", + ) + + assert scanned == ["mail a@b.com"] + assert result["texts"] == ["mail [EMAIL]"] + + @pytest.mark.asyncio async def test_request_without_user_text_is_not_scanned_or_recorded() -> None: respond, scanned = _engine()