diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 7a110eff080..a6864d93454 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -9986,7 +9986,7 @@ }, "unreachable_fallback": { "default": "fail_closed", - "description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.", + "description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'neuraltrust'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.", "enum": [ "fail_closed", "fail_open" @@ -11450,6 +11450,18 @@ "title": "Chunk Budget Chars", "type": "integer" }, + "collector_key": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "TrustGuard collector key (tgcol_...). Optional when the API key is bound to a collector. Env: TRUSTGUARD_COLLECTOR_KEY.", + "title": "Collector Key" + }, "confidence_threshold": { "default": 0.5, "default_value": 0.5, diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md new file mode 100644 index 00000000000..d6ba8f635bb --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md @@ -0,0 +1,60 @@ +# NeuralTrust TrustGuard + +Native LiteLLM guardrail. Sends chat input and output to TrustGuard `POST /v1/evaluate`. + +Setup guide, verdict mapping, and the streaming caveat: +[docs.neuraltrust.ai/integrations/litellm](https://docs.neuraltrust.ai/integrations/litellm). + +## Config + +```yaml +guardrails: + - guardrail_name: neuraltrust-trustguard + litellm_params: + guardrail: neuraltrust + mode: [pre_call, post_call] + api_key: os.environ/TRUSTGUARD_API_KEY + api_base: os.environ/TRUSTGUARD_API_BASE # default https://trustguard.neuraltrust.ai + collector_key: os.environ/TRUSTGUARD_COLLECTOR_KEY # tgcol_… ; optional if the API key is bound + unreachable_fallback: fail_closed + timeout: 5 + default_on: true +``` + +## Auth + +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 | +| --- | --- | +| `block` | HTTP 400 (trace_id / request_id only; findings are not echoed) | +| `ask` | HTTP 400 like `block`: a proxy has no approval flow, so the response names `verdict: ask` | +| `transform` | rewrite the last user message / last text from `transformed_payload` | +| `report` / `allow` | pass through (`report` is logged by trace_id) | + +Unknown verdicts, malformed bodies, and `transform` without a usable payload fail closed. + +## Fail-open vs fail-closed + +`unreachable_fallback` applies only to transport failures: connect errors, timeouts, HTTP 502/504. + +HTTP 503 entitlements, 401/403, other 4xx/5xx, and unusable TrustGuard verdicts always fail closed. + +`fail_open` means the request bypasses TrustGuard entirely when the endpoint is unreachable. It is off by default. + +## Streaming + +LiteLLM streaming guardrails default to `block_only`. `block` still fires on streamed calls. `transform` rewrites are not applied to the streamed tokens; use non-streaming requests when DLP redaction must reach the client. + +## References + +- [NeuralTrust TrustGuard on LiteLLM](https://docs.neuraltrust.ai/integrations/litellm) +- [TrustGuard Evaluate API](https://docs.neuraltrust.ai/trustguard/api/evaluate) +- [TrustGuard collectors](https://docs.neuraltrust.ai/trustguard/concepts/collectors) +- [LiteLLM Guardrails Documentation](https://docs.litellm.ai/docs/proxy/guardrails/quick_start) diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/__init__.py new file mode 100644 index 00000000000..5c11d1e173f --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/__init__.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Final + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .neuraltrust import NeuralTrustGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> NeuralTrustGuardrail: + import litellm + + _callback: Final = NeuralTrustGuardrail( + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + collector_key=litellm_params.collector_key, + unreachable_fallback=litellm_params.unreachable_fallback, + timeout=litellm_params.timeout, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(_callback) + return _callback + + +guardrail_initializer_registry: Final = { # mutable-ok: guardrail_registry discovers dict registries + SupportedGuardrailIntegrations.NEURALTRUST.value: initialize_guardrail, +} + +guardrail_class_registry: Final = { # mutable-ok: guardrail_registry discovers dict registries + SupportedGuardrailIntegrations.NEURALTRUST.value: NeuralTrustGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py new file mode 100644 index 00000000000..1106f856c5b --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py @@ -0,0 +1,406 @@ +"""NeuralTrust TrustGuard native LiteLLM guardrail. + +Calls TrustGuard POST /v1/evaluate on pre_call (input) and post_call (output). +""" + +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 +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + get_session_id_from_request_data, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.guardrails import GuardrailEventHooks, Mode +from litellm.types.proxy.guardrails.guardrail_hooks.neuraltrust import DEFAULT_API_BASE, DEFAULT_TIMEOUT +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel + +EVALUATE_PATH: Final = "/v1/evaluate" +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_ASK: Final = "ask" +STATUS_TRANSFORM: Final = "transform" +STATUS_REPORT: Final = "report" +STATUS_ALLOW: Final = "allow" +BLOCKING_STATUSES: Final = frozenset({STATUS_BLOCK, STATUS_ASK}) +KNOWN_STATUSES: Final = frozenset({STATUS_ALLOW, STATUS_TRANSFORM, STATUS_REPORT, *BLOCKING_STATUSES}) +UNREACHABLE_HTTP_STATUSES: Final = frozenset({502, 504}) +TRANSFORM_MISSING: Final = "TrustGuard transform missing payload" + + +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 "" + + +def _copy_message(value: object) -> Mapping[str, object] | None: + if not isinstance(value, Mapping): + return None + return {str(key): item for key, item in value.items()} # mutable-ok: shallow copy for write-back + + +def _copy_messages(messages: Sequence[object]) -> tuple[Mapping[str, object], ...] | None: + copied: Final = tuple(copy for message in messages if (copy := _copy_message(message)) is not None) + return copied if len(copied) == len(messages) else None + + +def _texts_from_messages(messages: Sequence[Mapping[str, object]]) -> tuple[str, ...]: + return tuple(_message_text(message) for message in messages) + + +def _tool_calls_in_message(message: Mapping[str, object]) -> tuple[object, ...] | None: + if "tool_calls" not in message: + return None + raw: Final = message["tool_calls"] + if not isinstance(raw, list): + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + return tuple(raw) + + +def _tool_calls_from_messages(messages: Sequence[Mapping[str, object]]) -> tuple[object, ...] | None: + groups: Final = tuple(_tool_calls_in_message(message) for message in messages) + if all(group is None for group in groups): + return None + return tuple(tool_call for group in groups if group is not None for tool_call in group) + + +def _rewrite_last_user_message( + messages: Sequence[Mapping[str, object]], + redacted: str, +) -> tuple[Mapping[str, object], ...]: + user_indices: Final = tuple(index for index, message in enumerate(messages) if message.get("role") == "user") + target: Final = user_indices[-1] if user_indices else len(messages) - 1 + if target < 0: + return ({"role": "user", "content": redacted},) # mutable-ok: write-back message + return tuple( + {**message, "content": redacted} if index == target else dict(message) # mutable-ok: write-back message + for index, message in enumerate(messages) + ) + + +def _model_name( + inputs: GenericGuardrailAPIInputs, + logging_obj: LiteLLMLoggingObj | None, +) -> str: + if logging_obj is not None and logging_obj.model: + return str(logging_obj.model) + return str(inputs.get("model") or "") + + +def _assistant_message(text: str | None, tool_calls: object) -> Mapping[str, object]: + if tool_calls: + return {"role": "assistant", "content": text, "tool_calls": tool_calls} # mutable-ok: outbound JSON + return {"role": "assistant", "content": text} # mutable-ok: outbound JSON + + +def _assistant_messages(texts: Sequence[str], tool_calls: object) -> tuple[Mapping[str, object], ...]: + if not texts: + return (_assistant_message(None if tool_calls else "", tool_calls),) + last: Final = len(texts) - 1 + return tuple(_assistant_message(text, tool_calls if index == last else None) for index, text in enumerate(texts)) + + +def _sent_messages( + inputs: GenericGuardrailAPIInputs, + input_type: Literal["request", "response"], +) -> Sequence[Mapping[str, object]]: + if input_type == "response": + return _assistant_messages(tuple(inputs.get("texts") or ()), inputs.get("tool_calls")) + structured: Final = inputs.get("structured_messages") + if structured: + return structured + return tuple({"role": "user", "content": text} for text in (inputs.get("texts") or ())) # mutable-ok: outbound JSON + + +def _inputs_with_messages( + inputs: GenericGuardrailAPIInputs, + messages: Sequence[Mapping[str, object]], + *, + replace_tool_calls: bool, +) -> GenericGuardrailAPIInputs: + extracted: Final = _tool_calls_from_messages(messages) if replace_tool_calls else None + original_tool_calls: Final = inputs.get("tool_calls") + if extracted is not None and original_tool_calls is not None and len(extracted) != len(original_tool_calls): + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + merged: Final[GenericGuardrailAPIInputs] = { # mutable-ok: GenericGuardrailAPIInputs is a TypedDict + **inputs, + "structured_messages": list(messages), # mutable-ok: GenericGuardrailAPIInputs.structured_messages is a list + } + rebuilt: Final[GenericGuardrailAPIInputs] = ( + {**merged, "texts": list(_texts_from_messages(messages))} # mutable-ok: TypedDict field is a list + if inputs.get("texts") + else merged + ) + if extracted is None: + return rebuilt + return {**rebuilt, "tool_calls": list(extracted)} # mutable-ok: GenericGuardrailAPIInputs.tool_calls is a list + + +class NeuralTrustGuardrail(CustomGuardrail): + """LiteLLM hook that evaluates prompts and completions with TrustGuard.""" + + @staticmethod + def get_config_model() -> type[GuardrailConfigModel]: + from litellm.types.proxy.guardrails.guardrail_hooks.neuraltrust import ( + NeuralTrustGuardrailConfigModel, + ) + + return NeuralTrustGuardrailConfigModel + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: CustomGuardrail contract + return [ # mutable-ok: CustomGuardrail.supported_event_hooks is a list + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ] + + def __init__( + self, + api_base: str | None = None, + api_key: str | None = None, + collector_key: str | None = None, + unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", + timeout: float | None = None, + guardrail_name: str | None = None, + event_hook: GuardrailEventHooks | Mode | str | Sequence[str] | None = None, + default_on: bool | None = None, + ) -> None: + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback, + ) + self.api_base = (api_base or os.environ.get("TRUSTGUARD_API_BASE") or DEFAULT_API_BASE).rstrip("/") + self.api_key = api_key or os.environ.get("TRUSTGUARD_API_KEY") or "" + if not self.api_key: + raise ValueError( + "TrustGuard API key is required. Set TRUSTGUARD_API_KEY or pass api_key in litellm_params." + ) + 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, + supported_event_hooks=self.get_supported_event_hooks(), + # LitellmParams.mode is str | list[str] | Mode, which CustomGuardrail narrows to the enum + event_hook=event_hook, # pyright: ignore[reportArgumentType] # config supplies the raw mode string + default_on=bool(default_on), + ) + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, # mutable-ok: CustomGuardrail.apply_guardrail contract + input_type: Literal["request", "response"], + logging_obj: LiteLLMLoggingObj | None = None, + ) -> GenericGuardrailAPIInputs: + body: Final = self._evaluate_body(inputs, request_data, input_type, logging_obj) + try: + result: Final = await self._call_evaluate(body) + except HTTPException: + raise + except _TrustGuardUnreachable as exc: + return self._handle_unreachable(inputs, exc) + + status: Final = result["status"] + if status in BLOCKING_STATUSES: + raise HTTPException( + status_code=400, + detail={ # mutable-ok: FastAPI HTTPException.detail is a JSON object + "error": "Violated guardrail policy", + "neuraltrust_guardrail_response": "Blocked by NeuralTrust TrustGuard.", + "verdict": status, + "trace_id": result.get("trace_id"), + "request_id": result.get("request_id"), + }, + ) + if status == STATUS_TRANSFORM: + return self._apply_transform( + inputs, + result.get("transformed_payload"), + sent_count=len(_sent_messages(inputs, input_type)), + ) + if status == STATUS_REPORT: + verbose_proxy_logger.info("TrustGuard report-only findings trace_id=%s", result.get("trace_id")) + return inputs + + def _evaluate_body( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, # mutable-ok: CustomGuardrail.apply_guardrail contract + input_type: Literal["request", "response"], + 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", + "protocol": "llm", + "attributes": { # mutable-ok: outbound JSON + "content_type": "application/json", + "model": {"name": _model_name(inputs, logging_obj)}, # mutable-ok: outbound JSON + }, + **({"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 + def _payload( + inputs: GenericGuardrailAPIInputs, + input_type: Literal["request", "response"], + ) -> Mapping[str, object]: + messages: Final = _sent_messages(inputs, input_type) + tools: Final = inputs.get("tools") if input_type == "request" else None + if tools: + return {"messages": messages, "tools": tools} # mutable-ok: outbound JSON + return {"messages": messages} # mutable-ok: outbound JSON + + async def _call_evaluate(self, body: dict[str, object]) -> dict[str, object]: # mutable-ok: TrustGuard JSON + url: Final = f"{self.api_base}{EVALUATE_PATH}" + headers: Final = { # mutable-ok: AsyncHTTPHandler.post declares headers as dict + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + } + try: + response: Final = await self.async_handler.post( + url, + json=body, + headers=headers, + timeout=self.timeout, + ) + response.raise_for_status() + except Timeout as exc: + raise _TrustGuardUnreachable(exc) from exc + except httpx.HTTPStatusError as exc: + status_code: Final = exc.response.status_code + if status_code == 503: + raise HTTPException( + status_code=503, + detail="TrustGuard entitlements unavailable", + ) from exc + if status_code in (401, 403): + raise HTTPException( + status_code=status_code, + detail="TrustGuard authentication failed", + ) from exc + if status_code in UNREACHABLE_HTTP_STATUSES: + raise _TrustGuardUnreachable(exc) from exc + raise HTTPException( + status_code=503, + detail="TrustGuard request failed", + ) from exc + except httpx.RequestError as exc: + raise _TrustGuardUnreachable(exc) from exc + + try: + parsed: Final[object] = response.json() + except ValueError as exc: + raise _TrustGuardUnreachable("TrustGuard returned non-JSON body") from exc + if not isinstance(parsed, dict): + raise HTTPException(status_code=503, detail="TrustGuard returned an invalid response") + status: Final = parsed.get("status") + if not isinstance(status, str) or status.lower() not in KNOWN_STATUSES: + raise HTTPException(status_code=503, detail="TrustGuard returned an unknown verdict") + return {**parsed, "status": status.lower()} # mutable-ok: TrustGuard JSON object + + def _handle_unreachable( + self, + inputs: GenericGuardrailAPIInputs, + error: Exception, + ) -> GenericGuardrailAPIInputs: + if self.unreachable_fallback == "fail_open": + verbose_proxy_logger.critical( + "TrustGuard unreachable (fail-open): %s", + error, + exc_info=error, + ) + return inputs + verbose_proxy_logger.error("TrustGuard unreachable (fail-closed): %s", error) + raise HTTPException( + status_code=503, + detail="TrustGuard guardrail service unreachable", + ) from error + + @staticmethod + def _apply_transform( + inputs: GenericGuardrailAPIInputs, + transformed: object, + *, + sent_count: int, + ) -> GenericGuardrailAPIInputs: + if not isinstance(transformed, Mapping): + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + + raw_messages: Final = transformed.get("messages") + if isinstance(raw_messages, list) and raw_messages: + rewritten_messages: Final = _copy_messages(raw_messages) + if rewritten_messages is None or len(rewritten_messages) != sent_count: + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + return _inputs_with_messages(inputs, rewritten_messages, replace_tool_calls=True) + + raw_input: Final = transformed.get("input") + if not isinstance(raw_input, str) or not raw_input: + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + + original_messages: Final = inputs.get("structured_messages") + if isinstance(original_messages, list) and original_messages: + copied: Final = _copy_messages(original_messages) + if copied is None: + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + return _inputs_with_messages( + inputs, + _rewrite_last_user_message(copied, raw_input), + replace_tool_calls=False, + ) + + original_texts: Final = tuple(inputs.get("texts") or ()) + if not original_texts: + raise HTTPException(status_code=400, detail=TRANSFORM_MISSING) + rewritten_texts: Final = (*original_texts[:-1], raw_input) + return {**inputs, "texts": list(rewritten_texts)} # mutable-ok: GenericGuardrailAPIInputs.texts is a list diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 69cb88bfa2f..5d2cec08b1b 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -38,6 +38,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, ) +from litellm.types.proxy.guardrails.guardrail_hooks.neuraltrust import ( + NeuralTrustGuardrailConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( OvalixGuardrailConfigModel, ) @@ -72,7 +75,7 @@ Pydantic object defining how to set guardrails on litellm proxy guardrails: - guardrail_name: "bedrock-pre-guard" litellm_params: - guardrail: bedrock # supported values: "akto", "aporia", "bedrock", "lakera", "zscaler_ai_guard" + guardrail: bedrock # supported values: "akto", "aporia", "bedrock", "lakera", "neuraltrust", "zscaler_ai_guard" mode: "during_call" guardrailIdentifier: ff6ujrregl1q guardrailVersion: "DRAFT" @@ -90,6 +93,7 @@ class SupportedGuardrailIntegrations(Enum): PRESIDIO = "presidio" HIDE_SECRETS = "hide-secrets" HIDDENLAYER = "hiddenlayer" + NEURALTRUST = "neuraltrust" AIM = "aim" CATO_NETWORKS = "cato_networks" PANGEA = "pangea" @@ -945,7 +949,7 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up default="fail_closed", description=( "Behavior when a guardrail endpoint is unreachable due to network errors. " - "Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. " + "Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'neuraltrust'. " "'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed." ), ) @@ -1080,6 +1084,7 @@ class LitellmParams( # pyright: ignore[reportIncompatibleVariableOverride] # o QualifireGuardrailConfigModel, BlockCodeExecutionGuardrailConfigModel, HiddenlayerGuardrailConfigModel, + NeuralTrustGuardrailConfigModel, QostodianNexusConfigModel, VigilGuardGuardrailConfigModel, SingulrGuardrailConfigModel, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py new file mode 100644 index 00000000000..b05e58106f5 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py @@ -0,0 +1,54 @@ +from typing import Final, Literal + +from pydantic import Field + +from .base import GuardrailConfigModel + +DEFAULT_API_BASE: Final = "https://trustguard.neuraltrust.ai" +DEFAULT_TIMEOUT: Final = 5.0 + + +class NeuralTrustGuardrailConfigModel(GuardrailConfigModel): + """Config for the NeuralTrust TrustGuard native LiteLLM hook.""" + + api_key: str | None = Field( + default=None, + description=("TrustGuard API key (tgk_...). If not provided, TRUSTGUARD_API_KEY is checked."), + ) + + api_base: str | None = Field( + default=None, + description=("TrustGuard API base URL. Default https://trustguard.neuraltrust.ai. Env: TRUSTGUARD_API_BASE."), + ) + + collector_key: str | None = Field( + default=None, + description=( + "TrustGuard collector key (tgcol_...). Optional when the API key is bound to a " + "collector. Env: TRUSTGUARD_COLLECTOR_KEY." + ), + ) + + unreachable_fallback: Literal["fail_closed", "fail_open"] = Field( + default="fail_closed", + description=( + "What to do on transport failures (connect errors, timeouts, HTTP 502/504). " + "'fail_closed' blocks the request; 'fail_open' allows it. " + "HTTP 503 entitlements, 401/403, other 4xx/5xx, unknown verdicts, and " + "unusable transform payloads always fail closed. " + "'fail_open' means the request bypasses TrustGuard entirely." + ), + ) + + 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 new file mode 100644 index 00000000000..3476243bdb5 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py @@ -0,0 +1,898 @@ +import os +from typing import Literal +from unittest.mock import AsyncMock, patch + +import httpx +import pytest +from fastapi import HTTPException +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 + + +def _response(payload: object, status_code: int = 200) -> Response: + request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate") + return Response(status_code, request=request, json=payload) + + +def _logging() -> LiteLLMLoggingObj: + return LiteLLMLoggingObj( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hello"}], + stream=False, + call_type="completion", + litellm_call_id="call-1", + function_id="fn-1", + start_time=None, + ) + + +def _guardrail( + *, + api_key: str = "tgk_test", + collector_key: str = "tgcol_test", + guardrail_name: str = "neuraltrust", + event_hook: str = "pre_call", + default_on: bool = False, + unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", + timeout: float | None = None, + api_base: str | None = None, +) -> NeuralTrustGuardrail: + return NeuralTrustGuardrail( + api_key=api_key, + collector_key=collector_key, + guardrail_name=guardrail_name, + event_hook=event_hook, + default_on=default_on, + unreachable_fallback=unreachable_fallback, + timeout=timeout, + api_base=api_base, + ) + + +class TestNeuralTrustGuardrail: + def setup_method(self) -> None: + for key in ("TRUSTGUARD_API_KEY", "TRUSTGUARD_API_BASE", "TRUSTGUARD_COLLECTOR_KEY"): + os.environ.pop(key, None) + + def teardown_method(self) -> None: + for key in ("TRUSTGUARD_API_KEY", "TRUSTGUARD_API_BASE", "TRUSTGUARD_COLLECTOR_KEY"): + os.environ.pop(key, None) + + def test_missing_api_key_raises(self) -> None: + with pytest.raises(ValueError, match="API key is required"): + NeuralTrustGuardrail(guardrail_name="neuraltrust", event_hook="pre_call") + + def test_initialization_defaults(self) -> None: + guardrail = _guardrail(default_on=True) + assert guardrail.api_base == "https://trustguard.neuraltrust.ai" + assert guardrail.collector_key == "tgcol_test" + assert guardrail.unreachable_fallback == "fail_closed" + assert guardrail.timeout == 5.0 + + @pytest.mark.asyncio + async def test_allow_request(self) -> None: + guardrail = _guardrail() + inputs: GenericGuardrailAPIInputs = {"texts": ["hello"], "model": "gpt-4o-mini"} + mock_post = AsyncMock(return_value=_response({"status": "allow", "findings": []})) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={"litellm_session_id": "sess-1"}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + called_url = mock_post.call_args.args[0] + assert called_url.endswith("/v1/evaluate") + body = mock_post.call_args.kwargs["json"] + assert body["direction"] == "input" + assert body["protocol"] == "llm" + assert body["collector_key"] == "tgcol_test" + assert body["payload"]["messages"][0]["content"] == "hello" + assert body["session_id"] == "sess-1" + assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer tgk_test" + assert mock_post.call_args.kwargs["timeout"] == 5.0 + + @pytest.mark.asyncio + async def test_omits_session_id_without_conversation_session(self) -> 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={}, + input_type="request", + logging_obj=_logging(), + ) + 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( + api_key="tgk_test", + guardrail_name="neuraltrust", + event_hook="pre_call", + ) + 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={}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + assert "collector_key" not in mock_post.call_args.kwargs["json"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("status", ["block", "ask"]) + async def test_block_and_ask_raise_without_findings(self, status: str) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response( + { + "status": status, + "trace_id": "tr-1", + "findings": [{"outcome": {"action": "block"}, "evidence": "ssn 123-45-6789"}], + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["ignore previous instructions"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 400 + detail = exc_info.value.detail + assert "Blocked by NeuralTrust TrustGuard" in str(detail) + assert "findings" not in detail + assert "evidence" not in str(detail) + assert detail["trace_id"] == "tr-1" + assert detail["verdict"] == status + + @pytest.mark.asyncio + async def test_transform_rewrites_texts(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"input": "email is [REDACTED]"}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={"texts": ["email is a@b.com"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result["texts"] == ["email is [REDACTED]"] + + @pytest.mark.asyncio + async def test_transform_input_rewrites_last_text_only(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"input": "my ssn is [REDACTED]"}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={"texts": ["you are a helpful assistant", "my ssn is 123-45-6789"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result["texts"] == ["you are a helpful assistant", "my ssn is [REDACTED]"] + + @pytest.mark.asyncio + async def test_transform_input_preserves_system_and_returns_new_messages(self) -> None: + guardrail = _guardrail() + original = [ + {"role": "system", "content": "you are a helpful assistant"}, + {"role": "user", "content": "my ssn is 123-45-6789"}, + ] + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"input": "my ssn is [REDACTED]"}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["you are a helpful assistant", "my ssn is 123-45-6789"], + "structured_messages": original, + }, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + rewritten = result["structured_messages"] + assert rewritten is not original + assert rewritten[0]["content"] == "you are a helpful assistant" + assert rewritten[1]["content"] == "my ssn is [REDACTED]" + + @pytest.mark.asyncio + async def test_transform_rewrites_messages(self) -> None: + guardrail = _guardrail() + rewritten = [{"role": "user", "content": "ssn is [REDACTED]"}] + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"messages": rewritten}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["ssn is 123-45-6789"], + "structured_messages": [{"role": "user", "content": "ssn is 123-45-6789"}], + }, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result["texts"] == ["ssn is [REDACTED]"] + assert result["structured_messages"] == rewritten + assert result["structured_messages"] is not rewritten + + @pytest.mark.asyncio + async def test_transform_messages_writes_back_tool_calls(self) -> None: + guardrail = _guardrail(event_hook="post_call") + original_tool_calls = [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": '{"ssn":"123-45-6789"}'}} + ] + rewritten_tool_calls = [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": '{"ssn":"[REDACTED]"}'}} + ] + rewritten = [{"role": "assistant", "content": None, "tool_calls": rewritten_tool_calls}] + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"messages": rewritten}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={ + "texts": [""], + "tool_calls": original_tool_calls, + "structured_messages": [{"role": "assistant", "content": None, "tool_calls": original_tool_calls}], + }, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert result["tool_calls"] == rewritten_tool_calls + assert result["tool_calls"] is not original_tool_calls + assert result["structured_messages"][0]["tool_calls"] == rewritten_tool_calls + + @pytest.mark.asyncio + @pytest.mark.parametrize("emptied", ["", None]) + async def test_transform_emptied_output_blanks_text_instead_of_restoring_original( + self, emptied: str | None + ) -> None: + guardrail = _guardrail(event_hook="post_call") + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"messages": [{"role": "assistant", "content": emptied}]}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={"texts": ["my ssn is 123-45-6789"]}, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert result["texts"] == [""] + + @pytest.mark.asyncio + async def test_transform_emptied_output_keeps_choice_alignment(self) -> None: + guardrail = _guardrail(event_hook="post_call") + rewritten = [ + {"role": "assistant", "content": ""}, + {"role": "assistant", "content": "card ending [REDACTED]"}, + ] + mock_post = AsyncMock( + return_value=_response({"status": "transform", "transformed_payload": {"messages": rewritten}}) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={"texts": ["ssn 123-45-6789", "card ending 4242"]}, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert result["texts"] == ["", "card ending [REDACTED]"] + + @pytest.mark.asyncio + async def test_transform_emptied_output_reaches_client_blank_and_aligned(self) -> None: + guardrail = _guardrail(event_hook="post_call") + rewritten = [ + {"role": "assistant", "content": ""}, + {"role": "assistant", "content": "card ending [REDACTED]"}, + ] + mock_post = AsyncMock( + return_value=_response({"status": "transform", "transformed_payload": {"messages": rewritten}}) + ) + response = ModelResponse( + id="chatcmpl-1", + created=1, + model="gpt-4o-mini", + object="chat.completion", + choices=[ + Choices(finish_reason="stop", index=0, message=Message(content="ssn 123-45-6789", role="assistant")), + Choices(finish_reason="stop", index=1, message=Message(content="card ending 4242", role="assistant")), + ], + ) + with patch.object(guardrail.async_handler, "post", mock_post): + processed = await OpenAIChatCompletionsHandler().process_output_response(response, guardrail) + assert processed.choices[0].message.content == "" + assert processed.choices[1].message.content == "card ending [REDACTED]" + + @pytest.mark.asyncio + @pytest.mark.parametrize("sent_texts", [{}, {"texts": []}]) + async def test_transform_tool_call_only_output_adds_no_text(self, sent_texts: GenericGuardrailAPIInputs) -> None: + guardrail = _guardrail(event_hook="post_call") + original_tool_calls = [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": '{"ssn":"123-45-6789"}'}} + ] + rewritten_tool_calls = [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": '{"ssn":"[REDACTED]"}'}} + ] + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": { + "messages": [{"role": "assistant", "content": None, "tool_calls": rewritten_tool_calls}] + }, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={**sent_texts, "tool_calls": original_tool_calls}, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert not result.get("texts") + assert result["tool_calls"] == rewritten_tool_calls + + @pytest.mark.asyncio + @pytest.mark.parametrize("emptied", ["", None]) + async def test_transform_emptied_input_blanks_text_and_message(self, emptied: str | None) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"messages": [{"role": "user", "content": emptied}]}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["my ssn is 123-45-6789"], + "structured_messages": [{"role": "user", "content": "my ssn is 123-45-6789"}], + }, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result["texts"] == [""] + assert result["structured_messages"] == [{"role": "user", "content": emptied}] + + @pytest.mark.asyncio + @pytest.mark.parametrize("returned", [1, 3]) + async def test_transform_output_message_count_mismatch_fail_closed(self, returned: int) -> None: + guardrail = _guardrail(event_hook="post_call") + rewritten = [{"role": "assistant", "content": "[REDACTED]"} for _ in range(returned)] + mock_post = AsyncMock( + return_value=_response({"status": "transform", "transformed_payload": {"messages": rewritten}}) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["ssn 111-11-1111", "ssn 222-22-2222"]}, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 400 + assert "transform missing payload" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_transform_input_message_count_mismatch_fail_closed(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"messages": [{"role": "user", "content": "ssn is [REDACTED]"}]}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={ + "texts": ["you are a helpful assistant", "ssn is 123-45-6789"], + "structured_messages": [ + {"role": "system", "content": "you are a helpful assistant"}, + {"role": "user", "content": "ssn is 123-45-6789"}, + ], + }, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_transform_messages_keeps_tool_calls_when_omitted(self) -> None: + guardrail = _guardrail() + original_tool_calls = [ + {"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": '{"q":"hi"}'}} + ] + rewritten = [{"role": "user", "content": "ssn is [REDACTED]"}] + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": {"messages": rewritten}, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs={ + "texts": ["ssn is 123-45-6789"], + "tool_calls": original_tool_calls, + "structured_messages": [{"role": "user", "content": "ssn is 123-45-6789"}], + }, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result["tool_calls"] is original_tool_calls + + @pytest.mark.asyncio + async def test_transform_messages_tool_call_count_mismatch_fail_closed(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response( + { + "status": "transform", + "transformed_payload": { + "messages": [ + { + "role": "assistant", + "content": None, + "tool_calls": [], + } + ] + }, + } + ) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={ + "texts": [""], + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + }, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 400 + assert "transform missing payload" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_post_call_attaches_tool_calls_to_last_assistant_message(self) -> None: + guardrail = _guardrail(event_hook="post_call") + tool_calls = [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": '{"q":"hi"}'}}] + mock_post = AsyncMock(return_value=_response({"status": "allow"})) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs={"texts": ["first", "second"], "tool_calls": tool_calls}, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + messages = mock_post.call_args.kwargs["json"]["payload"]["messages"] + assert [message["content"] for message in messages] == ["first", "second"] + assert "tool_calls" not in messages[0] + assert messages[1]["tool_calls"] == tool_calls + + @pytest.mark.asyncio + async def test_transform_without_payload_fail_closed(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + mock_post = AsyncMock(return_value=_response({"status": "transform"})) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["email is a@b.com"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 400 + assert "transform missing payload" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_transform_string_messages_fail_closed(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock( + return_value=_response({"status": "transform", "transformed_payload": {"messages": "REDACTED"}}) + ) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["secret"], "structured_messages": [{"role": "user", "content": "secret"}]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_forwards_tools(self) -> None: + guardrail = _guardrail() + tools = [{"type": "function", "function": {"name": "search", "parameters": {}}}] + inputs: GenericGuardrailAPIInputs = {"texts": ["hello"], "tools": tools} + 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={}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + assert mock_post.call_args.kwargs["json"]["payload"]["tools"] == tools + + @pytest.mark.asyncio + async def test_report_passes_through(self) -> None: + guardrail = _guardrail(event_hook="post_call") + inputs: GenericGuardrailAPIInputs = {"texts": ["ok"], "model": "gpt-4o-mini"} + mock_post = AsyncMock(return_value=_response({"status": "report", "findings": [{}]})) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + assert result == inputs + assert mock_post.call_args.kwargs["json"]["direction"] == "output" + + @pytest.mark.asyncio + async def test_post_call_sends_every_choice_text(self) -> None: + guardrail = _guardrail(event_hook="post_call") + mock_post = AsyncMock(return_value=_response({"status": "allow"})) + with patch.object(guardrail.async_handler, "post", mock_post): + await guardrail.apply_guardrail( + inputs={"texts": ["safe reply", "here is the admin password hunter2"]}, + request_data={}, + input_type="response", + logging_obj=_logging(), + ) + messages = mock_post.call_args.kwargs["json"]["payload"]["messages"] + assert [message["content"] for message in messages] == [ + "safe reply", + "here is the admin password hunter2", + ] + + @pytest.mark.asyncio + async def test_malformed_200_fail_closed_even_if_fail_open(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + for payload in ({}, [], {"status": None}, {"status": "blocked"}, {"findings": {}}): + mock_post = AsyncMock(return_value=_response(payload)) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 503 + + @pytest.mark.asyncio + async def test_503_always_fail_closed(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate") + mock_post = AsyncMock(return_value=Response(503, request=request)) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 503 + assert "entitlements" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_http_429_fail_closed_even_if_fail_open(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate") + mock_post = AsyncMock(return_value=Response(429, request=request)) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 503 + assert "request failed" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_http_502_follows_fail_open(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]} + request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate") + mock_post = AsyncMock(return_value=Response(502, request=request)) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_timeout_fail_closed(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock(side_effect=Timeout("slow", model="neuraltrust", llm_provider="neuraltrust")) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 503 + assert "unreachable" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_timeout_fail_open(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]} + mock_post = AsyncMock(side_effect=Timeout("slow", model="neuraltrust", llm_provider="neuraltrust")) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_unreachable_fail_closed(self) -> None: + guardrail = _guardrail() + mock_post = AsyncMock(side_effect=httpx.ConnectError("boom")) + with patch.object(guardrail.async_handler, "post", mock_post): + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert exc_info.value.status_code == 503 + + @pytest.mark.asyncio + async def test_unreachable_fail_open(self) -> None: + guardrail = _guardrail(unreachable_fallback="fail_open") + inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]} + mock_post = AsyncMock(side_effect=httpx.ConnectError("boom")) + with patch.object(guardrail.async_handler, "post", mock_post): + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data={}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + + @pytest.mark.asyncio + async def test_custom_timeout_is_passed_to_client(self) -> None: + guardrail = _guardrail(timeout=12) + 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={}, + input_type="request", + logging_obj=_logging(), + ) + assert result == inputs + assert mock_post.call_args.kwargs["timeout"] == 12.0 + + def test_get_config_model(self) -> None: + model = NeuralTrustGuardrail.get_config_model() + 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, + ) + from litellm.proxy.guardrails.guardrail_registry import ( + guardrail_class_registry, + guardrail_initializer_registry, + ) + + assert "neuraltrust" in guardrail_initializer_registry + assert guardrail_class_registry["neuraltrust"] is Registered diff --git a/ui/litellm-dashboard/public/assets/logos/neuraltrust.svg b/ui/litellm-dashboard/public/assets/logos/neuraltrust.svg new file mode 100644 index 00000000000..46a00fa2d3e --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/neuraltrust.svg @@ -0,0 +1,22 @@ + + + + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts index d0afc896260..2b7d827482c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts @@ -216,6 +216,12 @@ export const GUARDRAIL_PRESETS: Record = { mode: "pre_call", defaultOn: false, }, + neuraltrust: { + provider: "Neuraltrust", + guardrailNameSuggestion: "NeuralTrust TrustGuard", + mode: "pre_call", + defaultOn: false, + }, noma: { provider: "Noma", guardrailNameSuggestion: "Noma Security", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts index 9a9ab3a61d7..99fdfb6cb43 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -12,6 +12,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record = { panw: "palo_alto_networks.jpeg", cisco_ai_defense: "cisco.png", noma: "noma_security.png", + neuraltrust: "neuraltrust.svg", aporia: "aporia.png", aim: "aim_security.jpeg", cato_networks: "cato_networks.svg", @@ -53,4 +54,10 @@ describe("guardrail_garden_data logos", () => { expect(card.logo, `card ${card.id}`).not.toContain("/ui/assets/logos/"); } }); + + it("does not publish unsourced NeuralTrust eval numbers", () => { + const card = PARTNER_GUARDRAIL_CARDS.find((c) => c.id === "neuraltrust"); + expect(card).toBeDefined(); + expect(card?.eval).toBeUndefined(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index 165bd8f9967..0352d47417b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -319,6 +319,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ tags: ["Enterprise", "Security", "Prompt Injection", "PII"], providerKey: "CiscoAiDefense", }, + { + id: "neuraltrust", + name: "NeuralTrust", + description: + "TrustGuard runtime guardrails: prompt injection, toxicity, DLP, and policy enforcement on LLM input and output.", + category: "partner", + logo: guardrailLogoMap["NeuralTrust"], + tags: ["Security", "Prompt Injection", "DLP"], + providerKey: "Neuraltrust", + }, { id: "noma", name: "Noma Security", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx index c5e07fe9624..18533622e13 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx @@ -196,6 +196,20 @@ describe("guardrail_info_helpers", () => { expect(result.logo).toContain("noma_security.png"); }); + it("should resolve NeuralTrust logo and display name", () => { + populateGuardrailProviders({ + neuraltrust: { ui_friendly_name: "NeuralTrust" }, + }); + populateGuardrailProviderMap({ + neuraltrust: { ui_friendly_name: "NeuralTrust" }, + }); + + const result = getGuardrailLogoAndName("neuraltrust"); + + expect(result.displayName).toBe("NeuralTrust"); + expect(result.logo).toContain("neuraltrust.svg"); + }); + it("should resolve RepelloAI Argus logo and display name", () => { populateGuardrailProviders({ repelloai: { ui_friendly_name: "RepelloAI Argus" }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index fb3cf8f309a..34ebc5bb6ba 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -15,6 +15,7 @@ import lakeraAiLogo from "../../../../../public/assets/logos/lakeraai.jpeg"; import lassoLogo from "../../../../../public/assets/logos/lasso.png"; import litellmLogo from "../../../../../public/assets/logos/litellm_logo.jpg"; import microsoftAzureLogo from "../../../../../public/assets/logos/microsoft_azure.svg"; +import neuraltrustLogo from "../../../../../public/assets/logos/neuraltrust.svg"; import nomaSecurityLogo from "../../../../../public/assets/logos/noma_security.png"; import openaiSmallLogo from "../../../../../public/assets/logos/openai_small.svg"; import paloAltoNetworksLogo from "../../../../../public/assets/logos/palo_alto_networks.jpeg"; @@ -187,6 +188,7 @@ export const guardrailLogoMap = { "Aporia AI": aporiaLogo.src, "PANW Prisma AIRS": paloAltoNetworksLogo.src, "Cisco AI Defense": ciscoLogo.src, + NeuralTrust: neuraltrustLogo.src, "Noma Security": nomaSecurityLogo.src, "Javelin Guardrails": javelinLogo.src, "Pillar Guardrail": pillarLogo.src, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 7eadaa6c991..51b5c545f47 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16781,7 +16781,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) @@ -16887,7 +16886,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) @@ -23984,7 +23982,7 @@ export interface components { timeout?: number | null; /** * Unreachable Fallback - * @description Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed. + * @description Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'neuraltrust'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed. * @default fail_closed * @enum {string} */ @@ -30848,6 +30846,11 @@ export interface components { * @default 25000 */ chunk_budget_chars: number; + /** + * Collector Key + * @description TrustGuard collector key (tgcol_...). Optional when the API key is bound to a collector. Env: TRUSTGUARD_COLLECTOR_KEY. + */ + collector_key?: string | null; /** * Confidence Threshold * @description Only block or mask when detection confidence >= this value; below threshold, allow or log_only.