mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
refactor(guardrails): reuse secret and message text helpers for TrendAI
This commit is contained in:
parent
721d1efe03
commit
477a58a91f
3 changed files with 49 additions and 45 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue