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:
albertbausili 2026-09-13 13:25:10 +02:00
parent b778c2d412
commit 2d9b4a3eb8
4 changed files with 165 additions and 2 deletions

View file

@ -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 |

View file

@ -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

View file

@ -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"

View file

@ -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,