diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md index 5d94cb98d53..2563eef227e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md @@ -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 | diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py index 61a8ec8940d..d20cc3ce37c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py @@ -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 diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py index c4f3dd2c36a..b05e58106f5 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py @@ -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" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py index 7f325790059..458987a8cd1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py @@ -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,