mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(guardrails): expose the TrustGuard timeout in the UI and send the virtual key as consumer_id
The config model now declares timeout (default 5 seconds), so the Admin UI create and edit forms render a number input for it and /guardrails/ui/provider_specific_params advertises it. The shared LitellmParams.timeout keeps its None default for every other guardrail because BaseLitellmParams precedes this model in the MRO, and the hook rejects a non-positive value at startup Every evaluate call now carries consumer_id so TrustGuard Activity and per-consumer policies group by the LiteLLM key rather than by conversation. The identity is resolved tier by tier across both metadata blocks the proxy populates: key alias first, then the key's user email, user id, and team alias. The unified guardrail path seeds litellm_metadata with the alias under user_api_key_key_alias while the request metadata uses user_api_key_alias, so both names are accepted and only string values are ever sent Tests build request_data with the proxy's own helpers for the chat completions shape and the seeded-only shape used by MCP and pass-through, pin the fallback order, and cover the UI field set and the timeout defaults
This commit is contained in:
parent
b778c2d412
commit
2d9b4a3eb8
4 changed files with 165 additions and 2 deletions
|
|
@ -25,6 +25,10 @@ guardrails:
|
|||
|
||||
Bearer `tgk_…` API key. Address the collector with `collector_key`, or omit it when the key is already bound to one.
|
||||
|
||||
## Identity
|
||||
|
||||
Each evaluate call carries `session_id` from the LiteLLM session and `consumer_id` from the virtual key: the key alias, else the key's user email, user id, or team alias. TrustGuard Activity and per-consumer policies group by that value.
|
||||
|
||||
## Verdicts
|
||||
|
||||
| TrustGuard `status` | LiteLLM |
|
||||
|
|
|
|||
|
|
@ -7,10 +7,12 @@ from __future__ import annotations
|
|||
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import Timeout
|
||||
|
|
@ -24,7 +26,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.neuraltrust import DEFAULT_API_BASE
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.neuraltrust import DEFAULT_API_BASE, DEFAULT_TIMEOUT
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -32,7 +34,14 @@ if TYPE_CHECKING:
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
EVALUATE_PATH: Final = "/v1/evaluate"
|
||||
DEFAULT_TIMEOUT: Final = 5.0
|
||||
CONSUMER_ID_KEYS: Final = (
|
||||
("user_api_key_alias", "user_api_key_key_alias"),
|
||||
("user_api_key_user_email",),
|
||||
("user_api_key_user_id",),
|
||||
("user_api_key_team_alias",),
|
||||
)
|
||||
METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
EMPTY_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
STATUS_BLOCK: Final = "block"
|
||||
STATUS_TRANSFORM: Final = "transform"
|
||||
STATUS_REPORT: Final = "report"
|
||||
|
|
@ -46,6 +55,19 @@ class _TrustGuardUnreachable(Exception):
|
|||
"""Transport or availability failure; eligible for unreachable_fallback."""
|
||||
|
||||
|
||||
def _metadata(block: object) -> Mapping[str, object]:
|
||||
try:
|
||||
return METADATA_ADAPTER.validate_python(block)
|
||||
except ValidationError:
|
||||
return EMPTY_METADATA
|
||||
|
||||
|
||||
def _consumer_id(request_data: Mapping[str, object]) -> str | None:
|
||||
blocks: Final = tuple(_metadata(request_data.get(source)) for source in ("litellm_metadata", "metadata"))
|
||||
candidates: Final = (block.get(name) for names in CONSUMER_ID_KEYS for name in names for block in blocks)
|
||||
return next((value for value in candidates if isinstance(value, str) and value), None)
|
||||
|
||||
|
||||
def _message_text(message: Mapping[str, object]) -> str:
|
||||
content: Final = message.get("content")
|
||||
return content if isinstance(content, str) else ""
|
||||
|
|
@ -195,6 +217,8 @@ class NeuralTrustGuardrail(CustomGuardrail):
|
|||
self.collector_key = collector_key or os.environ.get("TRUSTGUARD_COLLECTOR_KEY") or ""
|
||||
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback
|
||||
resolved_timeout: Final = DEFAULT_TIMEOUT if timeout is None else float(timeout)
|
||||
if resolved_timeout <= 0:
|
||||
raise ValueError("TrustGuard timeout must be a positive number of seconds.")
|
||||
self.timeout = resolved_timeout
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
|
|
@ -249,6 +273,7 @@ class NeuralTrustGuardrail(CustomGuardrail):
|
|||
logging_obj: LiteLLMLoggingObj | None,
|
||||
) -> dict[str, object]: # mutable-ok: outbound JSON
|
||||
session_id: Final = get_session_id_from_request_data(request_data)
|
||||
consumer_id: Final = _consumer_id(request_data)
|
||||
return { # mutable-ok: outbound JSON
|
||||
"payload": self._payload(inputs, input_type),
|
||||
"direction": "input" if input_type == "request" else "output",
|
||||
|
|
@ -259,6 +284,7 @@ class NeuralTrustGuardrail(CustomGuardrail):
|
|||
},
|
||||
**({"collector_key": self.collector_key} if self.collector_key else {}), # mutable-ok: outbound JSON
|
||||
**({"session_id": session_id} if session_id else {}), # mutable-ok: outbound JSON
|
||||
**({"consumer_id": consumer_id} if consumer_id is not None else {}), # mutable-ok: outbound JSON
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from pydantic import Field
|
|||
from .base import GuardrailConfigModel
|
||||
|
||||
DEFAULT_API_BASE: Final = "https://trustguard.neuraltrust.ai"
|
||||
DEFAULT_TIMEOUT: Final = 5.0
|
||||
|
||||
|
||||
class NeuralTrustGuardrailConfigModel(GuardrailConfigModel):
|
||||
|
|
@ -39,6 +40,15 @@ class NeuralTrustGuardrailConfigModel(GuardrailConfigModel):
|
|||
),
|
||||
)
|
||||
|
||||
timeout: float | None = Field(
|
||||
default=DEFAULT_TIMEOUT,
|
||||
gt=0.0,
|
||||
description=(
|
||||
"Seconds to wait for each TrustGuard evaluate call before it counts as a "
|
||||
"transport failure and unreachable_fallback applies. Default 5."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "NeuralTrust"
|
||||
|
|
|
|||
|
|
@ -9,10 +9,15 @@ from httpx import Request, Response
|
|||
|
||||
from litellm.exceptions import Timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import get_provider_specific_params
|
||||
from litellm.proxy.guardrails.guardrail_hooks.neuraltrust.neuraltrust import (
|
||||
NeuralTrustGuardrail,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
from litellm.types.utils import Choices, GenericGuardrailAPIInputs, Message, ModelResponse
|
||||
|
||||
|
||||
|
|
@ -115,6 +120,97 @@ class TestNeuralTrustGuardrail:
|
|||
assert result == inputs
|
||||
assert "session_id" not in mock_post.call_args.kwargs["json"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("input_type", ["request", "response"])
|
||||
async def test_consumer_id_is_the_key_alias_on_proxy_shaped_request_data(
|
||||
self, input_type: Literal["request", "response"]
|
||||
) -> None:
|
||||
auth = UserAPIKeyAuth(key_alias="billing-app", user_id="u-1", user_email="dev@example.com", team_alias="team-x")
|
||||
request_data = {
|
||||
"metadata": LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=auth),
|
||||
"litellm_metadata": BaseTranslation.transform_user_api_key_dict_to_metadata(auth),
|
||||
}
|
||||
guardrail = _guardrail()
|
||||
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data=request_data,
|
||||
input_type=input_type,
|
||||
logging_obj=_logging(),
|
||||
)
|
||||
assert result == {"texts": ["hello"]}
|
||||
assert mock_post.call_args.kwargs["json"]["consumer_id"] == "billing-app"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_consumer_id_reads_the_seeded_key_alias_without_request_metadata(self) -> None:
|
||||
auth = UserAPIKeyAuth(key_alias="billing-app", user_email="dev@example.com")
|
||||
guardrail = _guardrail()
|
||||
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data={"litellm_metadata": BaseTranslation.transform_user_api_key_dict_to_metadata(auth)},
|
||||
input_type="request",
|
||||
logging_obj=_logging(),
|
||||
)
|
||||
assert result == {"texts": ["hello"]}
|
||||
assert mock_post.call_args.kwargs["json"]["consumer_id"] == "billing-app"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("request_data", "expected"),
|
||||
[
|
||||
(
|
||||
{"metadata": {"user_api_key_alias": "billing-app", "user_api_key_user_email": "dev@example.com"}},
|
||||
"billing-app",
|
||||
),
|
||||
(
|
||||
{"litellm_metadata": {"user_api_key_user_email": "dev@example.com", "user_api_key_user_id": "u-1"}},
|
||||
"dev@example.com",
|
||||
),
|
||||
({"metadata": {"user_api_key_user_id": 42, "user_api_key_team_alias": "team-x"}}, "team-x"),
|
||||
({"metadata": {"user_api_key_team_alias": "team-x"}}, "team-x"),
|
||||
(
|
||||
{
|
||||
"litellm_metadata": {"user_api_key_user_email": "dev@example.com"},
|
||||
"metadata": {"user_api_key_alias": "billing-app"},
|
||||
},
|
||||
"billing-app",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_consumer_id_falls_back_through_key_identity(self, request_data: dict, expected: str) -> None:
|
||||
guardrail = _guardrail()
|
||||
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hello"]},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
logging_obj=_logging(),
|
||||
)
|
||||
assert result == {"texts": ["hello"]}
|
||||
assert mock_post.call_args.kwargs["json"]["consumer_id"] == expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_data", [{}, {"metadata": {"user_api_key_alias": "", "user_api_key_user_id": None}}]
|
||||
)
|
||||
async def test_omits_consumer_id_without_key_identity(self, request_data: dict) -> None:
|
||||
guardrail = _guardrail()
|
||||
inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]}
|
||||
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
logging_obj=_logging(),
|
||||
)
|
||||
assert result == inputs
|
||||
assert "consumer_id" not in mock_post.call_args.kwargs["json"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_omits_collector_key_when_unbound(self) -> None:
|
||||
guardrail = NeuralTrustGuardrail(
|
||||
|
|
@ -760,6 +856,33 @@ class TestNeuralTrustGuardrail:
|
|||
assert model is not None
|
||||
assert model.ui_friendly_name() == "NeuralTrust"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ui_offers_timeout_with_the_connection_fields(self) -> None:
|
||||
fields = (await get_provider_specific_params())["neuraltrust"]
|
||||
assert fields["ui_friendly_name"] == "NeuralTrust"
|
||||
assert set(fields) - {"ui_friendly_name"} == {
|
||||
"api_key",
|
||||
"api_base",
|
||||
"collector_key",
|
||||
"unreachable_fallback",
|
||||
"timeout",
|
||||
}
|
||||
assert fields["timeout"]["type"] == "number"
|
||||
assert fields["timeout"]["default_value"] == 5.0
|
||||
assert fields["unreachable_fallback"]["options"] == ["fail_closed", "fail_open"]
|
||||
|
||||
def test_timeout_default_stays_local_to_neuraltrust(self) -> None:
|
||||
assert LitellmParams(guardrail="lakera_v2", mode="pre_call").timeout is None
|
||||
unset = LitellmParams(guardrail="neuraltrust", mode="pre_call").timeout
|
||||
explicit = LitellmParams(guardrail="neuraltrust", mode="pre_call", timeout=2).timeout
|
||||
assert _guardrail(timeout=unset).timeout == 5.0
|
||||
assert _guardrail(timeout=explicit).timeout == 2.0
|
||||
|
||||
@pytest.mark.parametrize("timeout", [0, -1.5])
|
||||
def test_rejects_non_positive_timeout(self, timeout: float) -> None:
|
||||
with pytest.raises(ValueError, match="positive"):
|
||||
_guardrail(timeout=timeout)
|
||||
|
||||
def test_registry_contains_neuraltrust(self) -> None:
|
||||
from litellm.proxy.guardrails.guardrail_hooks.neuraltrust import (
|
||||
NeuralTrustGuardrail as Registered,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue