diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml
index 2d399cca3a4..69d3ab25a1f 100644
--- a/.github/workflows/test-rust.yml
+++ b/.github/workflows/test-rust.yml
@@ -88,7 +88,7 @@ jobs:
rust-test:
runs-on: ubuntu-latest
- timeout-minutes: 20
+ timeout-minutes: 30
defaults:
run:
working-directory: litellm-rust
@@ -169,7 +169,11 @@ jobs:
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
- - run: python tests/unit/rust_bridge/native_route_wheel_test.py dist/*.whl
+ - name: Test native routes from the installed wheel
+ run: |
+ wheel=(dist/*.whl)
+ uv run --isolated --no-project --with "${wheel[0]}" \
+ python tests/unit/rust_bridge/native_route_wheel_test.py "${wheel[0]}"
- name: Run pytest tests/test_litellm_rust with the compiled extension
run: make test-rust-extension
diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py
index fa5512e7bfe..661cffa9846 100644
--- a/litellm/llms/openai/chat/guardrail_translation/handler.py
+++ b/litellm/llms/openai/chat/guardrail_translation/handler.py
@@ -18,9 +18,11 @@ import json
import time
import uuid
from collections.abc import Mapping, Sequence
+from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Union, cast
+from pydantic import TypeAdapter
from typing_extensions import NotRequired, ReadOnly, TypedDict
import litellm
@@ -46,7 +48,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
unappliable_request_rewrite,
)
from litellm.main import stream_chunk_builder
-from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
+from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk, ChatCompletionToolParam
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
coerce_stream_holdback_value,
)
@@ -556,7 +558,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
terminate the stream. Text rewrites are not propagated to the client here
(see ``_process_streaming_transform`` for the incremental_diff path) unless
``deliver_ended_stream_rewrites`` opts the ended-stream branch in."""
- has_stream_ended: Final = self._first_choice_has_finished(responses_so_far)
+ has_stream_ended: Final = deliver_ended_stream_rewrites or self._first_choice_has_finished(responses_so_far)
if has_stream_ended:
await self._process_ended_stream(
@@ -655,13 +657,36 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
model_response: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
- await self.process_output_response(
- response=model_response,
- guardrail_to_apply=guardrail_to_apply,
- litellm_logging_obj=litellm_logging_obj,
- user_api_key_dict=user_api_key_dict,
- request_data=request_data,
+ inspection_responses: Final = (
+ tuple(
+ model_response.model_copy(update=MappingProxyType({"choices": [choice]}))
+ for choice in model_response.choices
+ )
+ if pre_guardrail_tool_calls and len(model_response.choices) > 1
+ else (model_response,)
)
+ inspection_request_data: Final = request_data if request_data is not None else {}
+ try:
+ for inspection_response in inspection_responses:
+ inspection_request_data["response"] = inspection_response
+ inspection_request_data["responses"] = (
+ [
+ self._narrowed_to_choice(chunk, inspection_response.choices[0].index)
+ for chunk in responses_so_far
+ ]
+ if len(inspection_response.choices) == 1
+ else responses_so_far
+ )
+ await self.process_output_response(
+ response=inspection_response,
+ guardrail_to_apply=guardrail_to_apply,
+ litellm_logging_obj=litellm_logging_obj,
+ user_api_key_dict=user_api_key_dict,
+ request_data=inspection_request_data,
+ )
+ finally:
+ inspection_request_data["response"] = model_response
+ inspection_request_data["responses"] = responses_so_far
if not deliver_ended_stream_rewrites:
return
await self._write_ended_stream_text_rewrites(
@@ -794,6 +819,50 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
self.merge_user_api_key_metadata_into_request(request_data, user_api_key_dict)
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
+ if self._streamed_tool_call_fingerprints(responses_so_far):
+ assembled: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
+ request_data["response"] = assembled
+ if len(assembled.choices) > 1:
+ choice_rounds: Final = tuple(
+ (
+ choice,
+ StreamTransformSink(),
+ [self._narrowed_to_choice(chunk, choice.index) for chunk in responses_so_far],
+ )
+ for choice in assembled.choices
+ )
+ try:
+ for choice, choice_sink, choice_chunks in choice_rounds:
+ request_data["response"] = assembled.model_copy(update=MappingProxyType({"choices": [choice]}))
+ request_data["responses"] = choice_chunks
+ await self._process_streaming_transform(
+ responses_so_far=choice_chunks,
+ guardrail_to_apply=guardrail_to_apply,
+ litellm_logging_obj=litellm_logging_obj,
+ user_api_key_dict=user_api_key_dict,
+ request_data=request_data,
+ sink=choice_sink,
+ )
+ finally:
+ request_data["response"] = assembled
+ request_data["responses"] = responses_so_far
+ sink.mutated_text_per_choice = dict(
+ chain.from_iterable(
+ choice_sink.mutated_text_per_choice.items() for _, choice_sink, _ in choice_rounds
+ )
+ )
+ sink.holdback_per_choice = dict(
+ chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink, _ in choice_rounds)
+ )
+ return
+ tool_calls: Final = chain.from_iterable(choice.message.tool_calls or () for choice in assembled.choices)
+ inputs["tool_calls"] = TypeAdapter(list[ChatCompletionToolCallChunk]).validate_python(
+ tuple(
+ {"index": index, **converted}
+ for index, tool_call in enumerate(tool_calls)
+ if (converted := self._convert_tool_call_to_dict(tool_call)) is not None
+ )
+ )
if responses_so_far and getattr(responses_so_far[0], "model", None):
inputs["model"] = responses_so_far[0].model
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
@@ -1132,20 +1201,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
def _function_tool_call_fragments(
responses_so_far: Sequence["ModelResponseStream"],
) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]:
- """Group the stream's function tool-call fragments by their tool-call index, in
- the index order ``stream_chunk_builder`` lists the rebuilt tool calls, keeping
- only the indices the builder keeps (an id and a name somewhere in the stream)."""
fragments: Final = tuple(
- tool_call
+ (choice.index, tool_call)
for response in responses_so_far
for choice in response.choices
for tool_call in choice.delta.tool_calls or ()
if isinstance(tool_call, ChatCompletionDeltaToolCall)
)
- identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id)
- named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name)
+ identified: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.id)
+ named: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.function.name)
return tuple(
- tuple(fragment for fragment in fragments if fragment.index == index) for index in sorted(identified & named)
+ tuple(fragment for choice, fragment in fragments if (choice, fragment.index) == key)
+ for key in sorted(identified & named)
)
def _write_ended_stream_tool_call_rewrites(
@@ -1155,28 +1222,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
guardrail_name: str,
) -> None:
- """Write ended-stream guardrail tool-call rewrites back across the buffered
- chunks: the rewritten name and full arguments land in the tool call's first
- fragment and the arguments of its later fragments are blanked, mirroring the
- text write-back. A rewrite on a stream carrying more than one distinct choice
- index, or whose fragments do not line up with the rebuilt tool calls, is
- reported as undeliverable, so the pipeline executor discards it and releases
- the original chunks."""
post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response)
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
return
- stream_choice_indices: Final = frozenset(
- choice.index for response in responses_so_far for choice in response.choices
- )
fragments_by_tool_call: Final = self._function_tool_call_fragments(responses_so_far)
- if len(stream_choice_indices) != 1:
- from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
-
- raise UndeliverableStreamRewrite(
- guardrail_name,
- f"the stream carries {len(stream_choice_indices)} choices and tool-call rewrites are only written "
- "back on single-choice streams",
- )
if len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index ded4db6d2aa..1a6f812710e 100644
--- a/litellm/proxy/_lazy_openapi_snapshot.json
+++ b/litellm/proxy/_lazy_openapi_snapshot.json
@@ -10380,6 +10380,22 @@
"description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.",
"title": "Sticky Session Routing"
},
+ "streaming_transform_mode": {
+ "anyOf": [
+ {
+ "enum": [
+ "block_only",
+ "incremental_diff"
+ ],
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "description": "Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.",
+ "title": "Streaming Transform Mode"
+ },
"template_id": {
"anyOf": [
{
@@ -10406,7 +10422,7 @@
},
"unreachable_fallback": {
"default": "fail_closed",
- "description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. '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', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', 'typesafe', 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"
@@ -11905,6 +11921,18 @@
"description": "Client secret of the gateway's Entra app registration, used to perform the On-Behalf-Of exchange. Falls back to the AGENT365_CLIENT_SECRET environment variable.",
"title": "Client Secret"
},
+ "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,
@@ -13188,6 +13216,22 @@
"description": "When True (default), after sensitive data is detected and routed, all subsequent requests in the same session will continue routing to the same model.",
"title": "Sticky Session Routing"
},
+ "streaming_transform_mode": {
+ "anyOf": [
+ {
+ "enum": [
+ "block_only",
+ "incremental_diff"
+ ],
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "description": "Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.",
+ "title": "Streaming Transform Mode"
+ },
"template_id": {
"anyOf": [
{
diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py
index 6053ab26726..a58cdb834de 100644
--- a/litellm/proxy/guardrails/guardrail_endpoints.py
+++ b/litellm/proxy/guardrails/guardrail_endpoints.py
@@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, Union,
from urllib.parse import urlparse
from fastapi import APIRouter, Depends, HTTPException, Request
-from pydantic import BaseModel, ValidationError
+from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
@@ -2225,6 +2225,26 @@ async def test_custom_code_guardrail(
)
+_GUARDRAIL_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
+
+
+def _metadata_fields(value: object) -> Mapping[str, object]:
+ try:
+ return _GUARDRAIL_METADATA_ADAPTER.validate_python(value)
+ except ValidationError:
+ return {} # mutable-ok: empty fallback for the request_data dict contract
+
+
+def _guardrail_request_metadata(caller: object, proxy: object) -> Mapping[str, object]:
+ caller_fields: Final = tuple(
+ (key, value) for key, value in _metadata_fields(caller).items() if not key.startswith("user_api_key_")
+ )
+ key_identity: Final = tuple(
+ (key, value) for key, value in _metadata_fields(proxy).items() if key.startswith("user_api_key_")
+ )
+ return dict((*caller_fields, *key_identity)) # mutable-ok: request_data is the apply_guardrail dict contract
+
+
def _execution_timeout_response(timeout: float) -> TestCustomCodeGuardrailResponse:
return TestCustomCodeGuardrailResponse(
success=False,
@@ -2404,9 +2424,10 @@ async def apply_guardrail(
if litellm_logging_obj is not None:
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
+ metadata: Final = _guardrail_request_metadata(request.metadata, data.get("metadata"))
request_data: Final[dict] = {
**({"messages": request.messages} if request.messages is not None else {}),
- **({"metadata": request.metadata} if request.metadata is not None else {}),
+ **({"metadata": metadata} if request.metadata is not None or metadata else {}),
}
_input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type)
guardrailed_inputs: Final = await active_guardrail.apply_guardrail(
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..2f3136967de
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/README.md
@@ -0,0 +1,74 @@
+# 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
+ streaming_transform_mode: incremental_diff
+ 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
+
+`block` and `ask` end the stream in either mode. `streaming_transform_mode` decides whether a `transform` verdict reaches the client.
+
+| Mode | What the client receives |
+| --- | --- |
+| `block_only` (default) | the raw model tokens, so the redaction is dropped |
+| `incremental_diff` | TrustGuard's rewritten reply |
+
+Under `incremental_diff` the reply is held until the end-of-stream evaluate returns, so the first token arrives with the last. TrustGuard re-reads the whole reply on every scan and its redaction spans move as the reply grows, so releasing tokens early would let a later scan try to rewrite text already on the wire. Holding them also means a blocking verdict ends the stream with nothing sent at all.
+
+`incremental_diff` covers OpenAI chat completions streaming; other surfaces fall back to `block_only`.
+
+Under `incremental_diff`, LiteLLM buffers tool-call deltas until TrustGuard inspects the assembled response. A blocking verdict releases no tool-call arguments or finish signal. Tool calls retain their IDs and order, with transformed arguments written into the buffered deltas before delivery
+
+The default `block_only` mode still forwards tool-call deltas before inspection. Use `incremental_diff` or non-streaming requests when tool calls must be checked before delivery. The example selects `incremental_diff` for inspected text and tool arguments
+
+## 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..efd59a311d8
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/__init__.py
@@ -0,0 +1,37 @@
+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,
+ streaming_transform_mode=litellm_params.streaming_transform_mode,
+ 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..53fcc87d7ee
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/neuraltrust/neuraltrust.py
@@ -0,0 +1,440 @@
+"""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:
+ raw: Final = message.get("tool_calls")
+ if raw is None:
+ return None
+ 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
+ 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] = {
+ **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,
+ streaming_transform_mode: Literal["block_only", "incremental_diff"] | 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
+ # Read off the instance by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook.
+ self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = (
+ streaming_transform_mode or "block_only"
+ )
+ 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:
+ return self._with_holdback(
+ await self._evaluate(inputs, request_data, input_type, logging_obj),
+ input_type,
+ )
+
+ async def _evaluate(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict, # mutable-ok: CustomGuardrail.apply_guardrail contract
+ input_type: Literal["request", "response"],
+ logging_obj: LiteLLMLoggingObj | 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"), input_type=input_type)
+ if status == STATUS_REPORT:
+ verbose_proxy_logger.info("TrustGuard report-only findings trace_id=%s", result.get("trace_id"))
+ return inputs
+
+ def _with_holdback(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ input_type: Literal["request", "response"],
+ ) -> GenericGuardrailAPIInputs:
+ """Withhold the whole reply from each streaming round under incremental_diff.
+
+ TrustGuard re-reads the full reply every round and its redaction spans move as the reply grows, so a
+ later round can rewrite text an earlier one already streamed. The engine cannot retract streamed bytes
+ and answers that with stream_transform_underflow, so nothing is released until the end-of-stream round
+ forces the holdback to zero and flushes the final redacted reply.
+ """
+ texts: Final = inputs.get("texts")
+ if input_type == "request" or self.streaming_transform_mode != "incremental_diff" or not texts:
+ return inputs
+ held: Final[GenericGuardrailAPIInputs] = {
+ **inputs,
+ "stream_holdback_chars": [len(text) for text in texts], # mutable-ok: the field is a list
+ }
+ return held
+
+ 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,
+ *,
+ input_type: Literal["request", "response"],
+ ) -> 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) != len(_sent_messages(inputs, input_type)):
+ 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)
+
+ # structured_messages on a reply scan is the request conversation the framework attached for context,
+ # so a scalar payload rewrites the scanned reply text instead.
+ original_messages: Final = inputs.get("structured_messages") if input_type == "request" else None
+ 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/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py
index d68a55f9a88..4cc9ff37c73 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py
@@ -708,40 +708,10 @@ class UnifiedLLMGuardrails(CustomLogger):
try:
async for item in response:
- # v1 transforms only text. A chunk carrying tool_calls is passed
- # through raw so function-calling turns are not dropped, but ONLY
- # its tool-call fields are forwarded: content is stripped so any
- # response text (in the same delta, or in another choice of an n>1
- # chunk) can never bypass the transform. The original chunk is kept
- # in responses_so_far so its text is still accumulated + redacted +
- # emitted as synthetic deltas, and so the guardrail inspects the
- # assembled tool calls at end of stream (see the block inspection
- # below), matching block_only. finish_reason rides on the raw
- # tool-only chunk, so it is not recorded for the text flush.
if self._chunk_has_tool_calls(item):
saw_tool_calls = True
responses_so_far.append(item)
last_chunk = item
- # Fix #3 — flush accumulated text BEFORE the tool-call
- # passthrough. Without this, a stream of text chunks that
- # hasn't yet hit a sampled round can be trailed by a
- # tool-call chunk carrying finish_reason="tool_calls"; an
- # SSE-compliant client stops reading at that finish_reason
- # and drops the end-of-stream text flush that would follow.
- if saw_text_content:
- async for out in _round(item, is_final=False):
- yield out
- # Fix #1 — pass finish_reason_per_choice into the
- # passthrough so a mixed content+tool_call chunk defers its
- # finish_reason to the final text terminator (see the
- # _tool_call_passthrough_chunk docstring).
- tool_only = self._tool_call_passthrough_chunk(
- item,
- finish_reason_per_choice=finish_reason_per_choice,
- held_choices=_held_choices(held_chars_per_choice),
- )
- responses_yielded.append(tool_only)
- yield tool_only
continue
if self._is_trailing_metadata_chunk(item):
@@ -759,38 +729,62 @@ class UnifiedLLMGuardrails(CustomLogger):
# sampled round here would guardrail the same content twice.
if (
not end_of_stream_only
+ and not saw_tool_calls
and not self._chunk_has_finish_reason(item)
and chunk_counter % sampling_rate == 0
):
async for out in _round(item, is_final=False):
yield out
- # v1 does not transform streamed tool calls, but they must still go
- # through the guardrail's block decision. Run the block_only inspection
- # over the full assembled response so tool calls cannot bypass it.
- #
- # Pass a deep copy of responses_so_far — the block path routes through
- # ``_process_streaming_block_only`` which mutates ``delta.content``
- # in-place on the chunk objects it receives. For an n>1 chunk carrying
- # text on one choice and tool_calls (with finish_reason) on another,
- # ``has_stream_ended`` reads ``choices[0]`` alone and can miss the
- # terminal signal, letting the block path rewrite the raw accumulator.
- # The subsequent final ``_round`` would then re-read the already-mutated
- # text, producing double-application for a non-idempotent guardrail or a
- # ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow
- # list copy wouldn't help — the mutation is on the chunk objects
- # themselves — so we deepcopy.
if saw_tool_calls:
- async for out in self._inspect_full_response_for_block(
+ inspected_responses: Final = copy.deepcopy(responses_so_far)
+ async for out in self._inspect_full_response(
endpoint_translation=endpoint_translation,
guardrail_to_apply=guardrail_to_apply,
request_data=request_data,
user_api_key_dict=user_api_key_dict,
- responses_so_far=copy.deepcopy(responses_so_far),
+ responses_so_far=inspected_responses,
responses_yielded=responses_yielded,
):
yield out
+ if saw_text_content:
+ async for out in _round(last_chunk, is_final=False):
+ yield out
+ tool_chunks: Final = tuple(
+ self._tool_call_passthrough_chunk(
+ buffered_item,
+ finish_reason_per_choice=finish_reason_per_choice,
+ held_choices=_held_choices(held_chars_per_choice),
+ )
+ for buffered_item in inspected_responses
+ if self._chunk_has_tool_calls(buffered_item)
+ )
+
+ async def checked_tail() -> AsyncGenerator[object, None]:
+ try:
+ async for tail_chunk in self._emit_stream_tail(
+ last_chunk=last_chunk,
+ final_round=_round,
+ responses_so_far=responses_so_far,
+ responses_yielded=responses_yielded,
+ ):
+ yield tail_chunk
+ except _StreamTerminated as exc:
+ yield exc
+
+ tail_chunks: Final = tuple([chunk async for chunk in checked_tail()])
+ if tail_chunks and isinstance(tail_chunks[-1], _StreamTerminated):
+ for error_chunk in tail_chunks[:-1]:
+ yield error_chunk
+ return
+ for tool_only in tool_chunks:
+ responses_yielded.append(tool_only)
+ yield tool_only
+ for tail_chunk in tail_chunks:
+ yield tail_chunk
+ return
+
async for out in self._emit_stream_tail(
last_chunk=last_chunk,
final_round=_round,
@@ -818,7 +812,7 @@ class UnifiedLLMGuardrails(CustomLogger):
responses_yielded.append(trailing)
yield trailing
- async def _inspect_full_response_for_block(
+ async def _inspect_full_response(
self,
*,
endpoint_translation: _EndpointTranslation,
@@ -828,16 +822,8 @@ class UnifiedLLMGuardrails(CustomLogger):
responses_so_far: Sequence[object],
responses_yielded: Sequence[object],
) -> AsyncGenerator[object, None]:
- """Run the block-only guardrail inspection over the full assembled
- response (text + tool calls) so nothing bypasses the block decision.
-
- The guardrail's returned transforms are discarded here (v1 does not
- transform tool calls); only its block decision matters. A block is
- surfaced the same way as elsewhere: ModifyResponseException terminates the
- stream via the shared block handler; a GenericGuardrailAPI block raises and
- propagates, matching block_only.
- """
from litellm.integrations.custom_guardrail import ModifyResponseException
+ from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
try:
await endpoint_translation.process_output_streaming_response(
@@ -847,7 +833,10 @@ class UnifiedLLMGuardrails(CustomLogger):
user_api_key_dict=user_api_key_dict,
request_data=request_data,
stream_transform_sink=None,
+ deliver_ended_stream_rewrites=True,
)
+ except UndeliverableStreamRewrite as exc:
+ raise HTTPException(status_code=400, detail="Guardrail stream rewrite could not be applied") from exc
except ModifyResponseException as e:
if e.original_response is None:
e.original_response = responses_so_far
diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py
index 579a3f6322f..1b747321e1a 100644
--- a/litellm/types/guardrails.py
+++ b/litellm/types/guardrails.py
@@ -41,6 +41,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,
)
@@ -78,7 +81,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"
@@ -96,6 +99,7 @@ class SupportedGuardrailIntegrations(Enum):
PRESIDIO = "presidio"
HIDE_SECRETS = "hide-secrets"
HIDDENLAYER = "hiddenlayer"
+ NEURALTRUST = "neuraltrust"
AIM = "aim"
CATO_NETWORKS = "cato_networks"
PANGEA = "pangea"
@@ -1059,11 +1063,22 @@ 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', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. "
+ "Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', 'typesafe', and 'neuraltrust'. "
"'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed."
),
)
+ streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = Field(
+ default=None,
+ description=(
+ "Whether a guardrail's text rewrite reaches a streaming client. "
+ "Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the "
+ "same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a "
+ "block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks "
+ "and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only."
+ ),
+ )
+
extra_headers: list[str] | None = Field(
default=None,
description=(
@@ -1195,6 +1210,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..73f7db0299a
--- /dev/null
+++ b/litellm/types/proxy/guardrails/guardrail_hooks/neuraltrust.py
@@ -0,0 +1,65 @@
+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."
+ ),
+ )
+
+ streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = Field(
+ default=None,
+ description=(
+ "How a `transform` verdict reaches a streaming client. `block_only` (default) streams the raw "
+ "model chunks, so `block` and `ask` still end the stream but the redacted text is dropped. "
+ "`incremental_diff` withholds the model chunks and streams TrustGuard's rewritten text and tool arguments: "
+ "the reply arrives once the end-of-stream evaluate returns, and a blocking verdict ends the "
+ "stream with nothing already sent. OpenAI chat completions streaming only."
+ ),
+ )
+
+ @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..7a2c6fdfc35
--- /dev/null
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_neuraltrust.py
@@ -0,0 +1,1309 @@
+import json
+import os
+from collections.abc import AsyncIterator, Sequence
+from typing import Literal
+from unittest.mock import AsyncMock, patch
+
+import httpx
+import litellm
+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 import initialize_guardrail
+from litellm.proxy.guardrails.guardrail_hooks.neuraltrust.neuraltrust import (
+ NeuralTrustGuardrail,
+)
+from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails
+from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
+from litellm.types.guardrails import LitellmParams
+from litellm.types.utils import (
+ ChatCompletionDeltaToolCall,
+ Choices,
+ Delta,
+ Function,
+ GenericGuardrailAPIInputs,
+ Message,
+ ModelResponse,
+ ModelResponseStream,
+ StreamingChoices,
+)
+
+
+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,
+ streaming_transform_mode: Literal["block_only", "incremental_diff"] | 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,
+ streaming_transform_mode=streaming_transform_mode,
+ api_base=api_base,
+ )
+
+
+CARD_NUMBER = "4111 1111 1111 1111"
+REPLY_CHUNKS = (
+ "Here is ",
+ "the billing ",
+ "record. ",
+ "Card ",
+ "4111 1111 ",
+ "1111 1111 ",
+ "is on file.",
+)
+FULL_REPLY = "".join(REPLY_CHUNKS)
+REDACTED_REPLY = FULL_REPLY.replace(CARD_NUMBER, "[REDACTED]")
+
+
+def _stream_chunk(content: str, finish_reason: str | None = None) -> ModelResponseStream:
+ return ModelResponseStream(
+ model="gpt-4o-mini",
+ choices=[
+ StreamingChoices(index=0, delta=Delta(content=content, role="assistant"), finish_reason=finish_reason)
+ ],
+ )
+
+
+async def _upstream_reply() -> AsyncIterator[ModelResponseStream]:
+ for chunk in REPLY_CHUNKS:
+ yield _stream_chunk(chunk)
+ yield _stream_chunk("", finish_reason="stop")
+
+
+FORBIDDEN_TOOL = "wire_transfer"
+
+
+async def _upstream_tool_call(include_text: bool = True) -> AsyncIterator[ModelResponseStream]:
+ for chunk in REPLY_CHUNKS[:4] if include_text else ():
+ yield _stream_chunk(chunk)
+ yield ModelResponseStream(
+ model="gpt-4o-mini",
+ choices=[
+ StreamingChoices(
+ index=0,
+ delta=Delta(
+ content=None,
+ role="assistant",
+ tool_calls=[
+ ChatCompletionDeltaToolCall(
+ id="call_1",
+ type="function",
+ index=0,
+ function=Function(name=FORBIDDEN_TOOL, arguments='{"amount": 100000}'),
+ )
+ ],
+ ),
+ finish_reason="tool_calls",
+ )
+ ],
+ )
+
+
+def _tool_call_blocking_trustguard() -> AsyncMock:
+ async def _post(*_args: object, **kwargs: object) -> Response:
+ seen = json.dumps(kwargs["json"]["payload"])
+ if FORBIDDEN_TOOL not in seen:
+ return _response({"status": "allow"})
+ return _response({"status": "block", "trace_id": "trace-1"})
+
+ return AsyncMock(side_effect=_post)
+
+
+def _finish_reasons(items: Sequence[object]) -> list[str]:
+ return [
+ choice.finish_reason
+ for item in items
+ if isinstance(item, ModelResponseStream)
+ for choice in (item.choices or [])
+ if choice.finish_reason
+ ]
+
+
+def _redacting_trustguard(*, on_card: str = "transform") -> AsyncMock:
+ """A TrustGuard that only reacts once the whole card number is in the accumulated reply."""
+
+ async def _post(*_args: object, **kwargs: object) -> Response:
+ body = kwargs["json"]
+ seen = "".join(message["content"] or "" for message in body["payload"]["messages"])
+ if CARD_NUMBER not in seen:
+ return _response({"status": "allow"})
+ if on_card != "transform":
+ return _response({"status": on_card, "trace_id": "trace-1"})
+ return _response(
+ {
+ "status": "transform",
+ "transformed_payload": {
+ "messages": [{"role": "assistant", "content": seen.replace(CARD_NUMBER, "[REDACTED]")}]
+ },
+ }
+ )
+
+ return AsyncMock(side_effect=_post)
+
+
+def _guardrail_stream(
+ guardrail: NeuralTrustGuardrail,
+ upstream: AsyncIterator[ModelResponseStream] | None = None,
+) -> AsyncIterator[object]:
+ return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
+ user_api_key_dict=UserAPIKeyAuth(api_key="tgk_test", request_route="/v1/chat/completions"),
+ response=_upstream_reply() if upstream is None else upstream,
+ request_data={
+ "guardrail_to_apply": guardrail,
+ "model": "gpt-4o-mini",
+ "messages": [{"role": "user", "content": "what card is on file?"}],
+ },
+ )
+
+
+async def _drain_into(stream: AsyncIterator[object], sink: list[object]) -> None:
+ async for item in stream:
+ sink.append(item)
+
+
+def _deltas(items: Sequence[object]) -> list[str]:
+ return [
+ item.choices[0].delta.content
+ for item in items
+ if isinstance(item, ModelResponseStream) and item.choices and item.choices[0].delta.content
+ ]
+
+
+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_transform_messages_with_non_object_entry_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_transform_input_with_non_object_original_message_fail_closed(self) -> None:
+ guardrail = _guardrail()
+ mock_post = AsyncMock(
+ return_value=_response({"status": "transform", "transformed_payload": {"input": "[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": ["secret"]}, # pyright: ignore[reportArgumentType] # malformed on purpose
+ request_data={},
+ input_type="request",
+ logging_obj=_logging(),
+ )
+ assert exc_info.value.status_code == 400
+
+ @pytest.mark.asyncio
+ async def test_transform_input_without_any_text_fail_closed(self) -> None:
+ guardrail = _guardrail()
+ mock_post = AsyncMock(
+ return_value=_response({"status": "transform", "transformed_payload": {"input": "[REDACTED]"}})
+ )
+ 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": {}}]},
+ 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_null_tool_calls_keeps_the_original_ones(self) -> None:
+ guardrail = _guardrail(event_hook="post_call")
+ original_tool_calls = [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}]
+ mock_post = AsyncMock(
+ return_value=_response(
+ {
+ "status": "transform",
+ "transformed_payload": {"messages": [{"role": "assistant", "content": "ok", "tool_calls": None}]},
+ }
+ )
+ )
+ with patch.object(guardrail.async_handler, "post", mock_post):
+ result = await guardrail.apply_guardrail(
+ inputs={"texts": ["secret"], "tool_calls": original_tool_calls},
+ request_data={},
+ input_type="response",
+ logging_obj=_logging(),
+ )
+ assert result["texts"] == ["ok"]
+ assert result["tool_calls"] is original_tool_calls
+
+ @pytest.mark.asyncio
+ async def test_transform_non_list_tool_calls_fail_closed(self) -> None:
+ guardrail = _guardrail(event_hook="post_call")
+ mock_post = AsyncMock(
+ return_value=_response(
+ {
+ "status": "transform",
+ "transformed_payload": {"messages": [{"role": "assistant", "content": "ok", "tool_calls": {}}]},
+ }
+ )
+ )
+ with patch.object(guardrail.async_handler, "post", mock_post):
+ with pytest.raises(HTTPException) as exc_info:
+ await guardrail.apply_guardrail(
+ inputs={"texts": ["secret"], "tool_calls": [{"id": "call_1", "type": "function", "function": {}}]},
+ request_data={},
+ input_type="response",
+ 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
+ @pytest.mark.parametrize("status_code", [401, 403])
+ async def test_auth_failures_return_their_status_even_if_fail_open(self, status_code: int) -> None:
+ guardrail = _guardrail(unreachable_fallback="fail_open")
+ mock_post = AsyncMock(return_value=_response({"error": "bad key"}, status_code=status_code))
+ 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 == status_code
+ assert "authentication failed" in str(exc_info.value.detail)
+
+ @pytest.mark.asyncio
+ async def test_non_json_200_fail_closed(self) -> None:
+ guardrail = _guardrail()
+ request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate")
+ mock_post = AsyncMock(return_value=Response(200, request=request, text="captive portal"))
+ 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_non_json_200_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(200, request=request, text="captive portal"))
+ 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_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
+
+ @pytest.mark.asyncio
+ async def test_incremental_diff_redacts_a_value_split_across_streamed_chunks(self) -> None:
+ """The card straddles a sampled scan, so the earlier half must never leave before the redaction lands."""
+ guardrail = _guardrail(event_hook="post_call", default_on=True, streaming_transform_mode="incremental_diff")
+ with patch.object(guardrail.async_handler, "post", _redacting_trustguard()):
+ out = [item async for item in _guardrail_stream(guardrail)]
+ assert _deltas(out) == [REDACTED_REPLY]
+
+ @pytest.mark.asyncio
+ async def test_block_mid_stream_under_incremental_diff_sends_nothing(self) -> None:
+ guardrail = _guardrail(event_hook="post_call", default_on=True, streaming_transform_mode="incremental_diff")
+ received: list[object] = [] # mutable-ok: collects what the client saw before the block
+ with patch.object(guardrail.async_handler, "post", _redacting_trustguard(on_card="block")):
+ with pytest.raises(HTTPException) as exc_info:
+ await _drain_into(_guardrail_stream(guardrail), received)
+ assert exc_info.value.status_code == 400
+ assert exc_info.value.detail["verdict"] == "block"
+ assert _deltas(received) == []
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("include_text", [True, False])
+ async def test_tool_call_block_under_incremental_diff_sends_nothing(self, include_text: bool) -> None:
+ guardrail = _guardrail(event_hook="post_call", default_on=True, streaming_transform_mode="incremental_diff")
+ received: list[object] = [] # mutable-ok: collects what the client saw before the block
+ with patch.object(guardrail.async_handler, "post", _tool_call_blocking_trustguard()):
+ with pytest.raises(HTTPException) as exc_info:
+ await _drain_into(_guardrail_stream(guardrail, _upstream_tool_call(include_text)), received)
+ assert exc_info.value.status_code == 400
+ assert exc_info.value.detail["verdict"] == "block"
+ assert _deltas(received) == []
+ assert _finish_reasons(received) == []
+ assert received == []
+
+ @pytest.mark.asyncio
+ async def test_default_streaming_mode_leaves_the_transform_off_the_wire(self) -> None:
+ """Default stays block_only, where the framework streams the raw model chunks and drops rewrites."""
+ guardrail = _guardrail(event_hook="post_call", default_on=True)
+ assert guardrail.streaming_transform_mode == "block_only"
+ with patch.object(guardrail.async_handler, "post", _redacting_trustguard()):
+ out = [item async for item in _guardrail_stream(guardrail)]
+ streamed = "".join(_deltas(out))
+ assert CARD_NUMBER in streamed
+ assert "[REDACTED]" not in streamed
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ ("mode", "input_type", "expected"),
+ [
+ ("incremental_diff", "response", [len(REDACTED_REPLY)]),
+ ("incremental_diff", "request", None),
+ ("block_only", "response", None),
+ ],
+ )
+ async def test_holdback_covers_streamed_reply_scans_only(
+ self,
+ mode: Literal["block_only", "incremental_diff"],
+ input_type: Literal["request", "response"],
+ expected: list[int] | None,
+ ) -> None:
+ guardrail = _guardrail(event_hook="post_call", streaming_transform_mode=mode)
+ mock_post = AsyncMock(
+ return_value=_response(
+ {
+ "status": "transform",
+ "transformed_payload": {"messages": [{"role": "assistant", "content": REDACTED_REPLY}]},
+ }
+ )
+ )
+ with patch.object(guardrail.async_handler, "post", mock_post):
+ result = await guardrail.apply_guardrail(
+ inputs={"texts": [FULL_REPLY]},
+ request_data={},
+ input_type=input_type,
+ logging_obj=_logging(),
+ )
+ assert result.get("stream_holdback_chars") == expected
+
+ @pytest.mark.asyncio
+ async def test_transform_input_on_a_reply_rewrites_the_reply_not_the_prompt(self) -> None:
+ """The framework attaches the request conversation to reply scans; a scalar payload must ignore it."""
+ guardrail = _guardrail(event_hook="post_call")
+ mock_post = AsyncMock(
+ return_value=_response({"status": "transform", "transformed_payload": {"input": "card [REDACTED]"}})
+ )
+ with patch.object(guardrail.async_handler, "post", mock_post):
+ result = await guardrail.apply_guardrail(
+ inputs={
+ "texts": ["card 4111 1111 1111 1111"],
+ "structured_messages": [
+ {"role": "user", "content": "what card is on file?"},
+ {"role": "assistant", "content": "card 4111 1111 1111 1111"},
+ ],
+ },
+ request_data={},
+ input_type="response",
+ logging_obj=_logging(),
+ )
+ assert result["texts"] == ["card [REDACTED]"]
+
+ 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",
+ "streaming_transform_mode",
+ }
+ assert fields["timeout"]["type"] == "number"
+ assert fields["timeout"]["default_value"] == 5.0
+ assert fields["unreachable_fallback"]["options"] == ["fail_closed", "fail_open"]
+ assert fields["streaming_transform_mode"]["options"] == ["block_only", "incremental_diff"]
+
+ 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_initializer_wires_params_and_registers_the_callback(self) -> None:
+ params = LitellmParams(
+ guardrail="neuraltrust",
+ mode="post_call",
+ api_key="tgk_from_params",
+ api_base="https://trustguard.example.test/",
+ collector_key="tgcol_from_params",
+ unreachable_fallback="fail_open",
+ timeout=2,
+ streaming_transform_mode="incremental_diff",
+ default_on=True,
+ )
+ hook = initialize_guardrail(params, {"guardrail_name": "tg-prod"})
+ try:
+ assert hook.api_key == "tgk_from_params"
+ assert hook.api_base == "https://trustguard.example.test"
+ assert hook.collector_key == "tgcol_from_params"
+ assert hook.unreachable_fallback == "fail_open"
+ assert hook.timeout == 2.0
+ assert hook.streaming_transform_mode == "incremental_diff"
+ assert hook.guardrail_name == "tg-prod"
+ assert hook.default_on is True
+ assert hook in litellm.callbacks
+ finally:
+ litellm.logging_callback_manager.remove_callback_from_all_lists(hook)
+
+ 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/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py
index c90f88ec110..52a628412cd 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py
@@ -3,6 +3,9 @@
import logging
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Literal
+from collections.abc import AsyncIterator
+
+from pydantic import TypeAdapter
import pytest
@@ -41,7 +44,10 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import ResponsesAPIResponse
-from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices
+from litellm.types.utils import (
+ CallTypes, ChatCompletionMessageToolCall, Delta, GenericGuardrailAPIInputs,
+ ModelResponseStream, StreamingChoices,
+)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@@ -939,6 +945,42 @@ def _delta_text(item):
return item.choices[0].delta.content or ""
+class _ToolRedactingGuardrail(CustomGuardrail):
+ def __init__(self, require_tool_context: bool = False) -> None:
+ super().__init__(guardrail_name="tool-redactor", event_hook=GuardrailEventHooks.post_call, default_on=True)
+ self.streaming_transform_mode = "incremental_diff"
+ self.streaming_sampling_rate = 1
+ self.require_tool_context = require_tool_context
+
+ async def apply_guardrail(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict[str, object],
+ input_type: Literal["request", "response"],
+ logging_obj: object | None = None,
+ ) -> GenericGuardrailAPIInputs:
+ calls: Final = TypeAdapter(tuple[ChatCompletionMessageToolCall, ...]).validate_python(
+ inputs.get("tool_calls", ())
+ )
+ input_texts: Final = inputs.get("texts", ())
+ texts: Final = tuple(
+ "checked:" + text.replace("SECRET", "MASKED")
+ if not self.require_tool_context or (calls and index == len(input_texts) - 1) else text
+ for index, text in enumerate(input_texts)
+ )
+ return {
+ **inputs,
+ "texts": list(texts),
+ "stream_holdback_chars": [len(text) for text in texts],
+ "tool_calls": [
+ call.model_copy(update={"function": call.function.model_copy(update={
+ "arguments": call.function.arguments.replace("SECRET", "MASKED"),
+ })})
+ for call in calls
+ ],
+ }
+
+
class TestStreamingTransform:
"""Streaming text-transformation (incremental_diff) path on the OpenAI chat
completions streaming surface."""
@@ -947,6 +989,103 @@ class TestStreamingTransform:
def _use_openai_handler_mapping(self, monkeypatch):
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("require_tool_context", [False, True])
+ @pytest.mark.parametrize("include_text", [False, True])
+ @pytest.mark.parametrize("tool_count", [1, 2])
+ async def test_buffered_tool_arguments_are_rewritten_before_delivery(
+ self, include_text: bool, tool_count: int, require_tool_context: bool
+ ) -> None:
+ chunks: Final = (
+ *([_stream_chunk("hello SECRET")] if include_text else []),
+ ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[
+ {"index": index, "id": f"call_{index}", "type": "function",
+ "function": {"name": "contact", "arguments": '{"contact":"SEC'}}
+ for index in range(tool_count)
+ ]))]),
+ ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[
+ {"index": index, "function": {"arguments": 'RET"}'}} for index in range(tool_count)
+ ]), finish_reason="tool_calls")]),
+ ModelResponseStream(choices=[], usage={"prompt_tokens": 7, "completion_tokens": 11, "total_tokens": 18}),
+ )
+ out: Final = await _drive_stream(
+ UnifiedLLMGuardrails(), _ToolRedactingGuardrail(require_tool_context=require_tool_context), chunks
+ )
+ calls: Final = tuple(
+ call for chunk in out for choice in chunk.choices for call in choice.delta.tool_calls or ()
+ )
+ for index in range(tool_count):
+ arguments: Final = "".join(call.function.arguments or "" for call in calls if call.index == index)
+ assert arguments == '{"contact":"MASKED"}'
+ assert next(call.id for call in calls if call.index == index and call.id) == f"call_{index}"
+ assert "".join(_delta_text(chunk) for chunk in out) == ("checked:hello MASKED" if include_text else "")
+ assert any(choice.finish_reason == "tool_calls" for chunk in out for choice in chunk.choices)
+ assert out[-1].usage.total_tokens == 18
+ assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize("tool_choice_index", [0, 1])
+ async def test_tool_dependent_text_rewrites_keep_choice_context(self, tool_choice_index: int) -> None:
+ chunks: Final = (
+ ModelResponseStream(choices=[StreamingChoices(
+ index=index, delta=Delta(content="SECRET" if index == tool_choice_index else "plain"),
+ ) for index in (1, 0)]),
+ ModelResponseStream(choices=[StreamingChoices(
+ index=tool_choice_index,
+ delta=Delta(tool_calls=[{
+ "index": 0, "id": "call_context", "type": "function",
+ "function": {"name": "contact", "arguments": '{"contact":"SECRET"}'},
+ }]), finish_reason="tool_calls",
+ )]),
+ )
+ out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(True), chunks)
+ for index in (0, 1):
+ text: Final = "".join(
+ choice.delta.content or "" for chunk in out for choice in chunk.choices if choice.index == index
+ )
+ assert text == ("checked:MASKED" if index == tool_choice_index else "plain")
+ assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
+
+ @pytest.mark.asyncio
+ async def test_tool_rewrites_keep_completion_choices_separate(self) -> None:
+ async def response() -> AsyncIterator[ModelResponseStream]:
+ yield ModelResponseStream(choices=[StreamingChoices(
+ index=index,
+ delta=Delta(tool_calls=[{"index": 0, "id": f"call_{index}", "type": "function",
+ "function": {"name": "contact", "arguments": '{"contact":"SECRET"}'}}]),
+ finish_reason="tool_calls",
+ ) for index in range(2)])
+
+ iterator: Final = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
+ user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"),
+ response=response(), request_data={"guardrail_to_apply": _ToolRedactingGuardrail()},
+ )
+ chunks: Final = tuple([chunk async for chunk in iterator])
+ calls: Final = tuple(
+ (choice.index, call.id, call.function.arguments)
+ for chunk in chunks for choice in chunk.choices for call in choice.delta.tool_calls or ()
+ )
+ assert calls == ((0, "call_0", '{"contact":"MASKED"}'), (1, "call_1", '{"contact":"MASKED"}'))
+
+ @pytest.mark.asyncio
+ async def test_undeliverable_rewrite_is_a_closed_failure(self) -> None:
+ from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
+
+ class UndeliverableTranslation(OpenAIChatCompletionsHandler):
+ async def process_output_streaming_response(
+ self, *args: object, **kwargs: object
+ ) -> list[ModelResponseStream]:
+ raise UndeliverableStreamRewrite("tool-redactor", "unmappable tool fragments")
+
+ iterator: Final = UnifiedLLMGuardrails()._inspect_full_response(
+ endpoint_translation=UndeliverableTranslation(), guardrail_to_apply=_ToolRedactingGuardrail(),
+ request_data={}, user_api_key_dict=UserAPIKeyAuth(), responses_so_far=(), responses_yielded=(),
+ )
+ with pytest.raises(unified_module.HTTPException) as error:
+ await anext(iterator)
+ assert error.value.status_code == 400
+ assert error.value.detail == "Guardrail stream rewrite could not be applied"
+
@pytest.mark.asyncio
async def test_block_only_drops_text_rewrites(self):
"""Default block_only: the guardrail's uppercasing never reaches the
@@ -1367,14 +1506,15 @@ class TestStreamingTransform:
assert out[-2].choices[0].finish_reason == "stop"
@pytest.mark.asyncio
- async def test_tool_call_blocking_guardrail_is_enforced(self):
+ @pytest.mark.parametrize(("content", "allowed_scans"), [(None, 0), ("proposal", 0), ("proposal", 1)])
+ async def test_tool_call_blocking_guardrail_is_enforced(self, content: str | None, allowed_scans: int):
"""A guardrail that blocks on tool calls must terminate the incremental_diff
stream: tool calls go through the block decision, not bypass it."""
from litellm.exceptions import GuardrailRaisedException
class _ToolCallBlocker(_StreamingTextGuardrail):
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
- if input_type == "response" and inputs.get("tool_calls"):
+ if input_type == "response" and self.response_calls >= allowed_scans:
raise GuardrailRaisedException(
guardrail_name="tc-block",
message="blocked tool call",
@@ -1387,7 +1527,7 @@ class TestStreamingTransform:
StreamingChoices(
index=0,
delta=Delta(
- content=None,
+ content=content,
tool_calls=[
{
"index": 0,
@@ -1402,8 +1542,21 @@ class TestStreamingTransform:
],
)
+ async def upstream():
+ yield tool_chunk
+
+ stream: Final = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
+ user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"),
+ response=upstream(),
+ request_data={"guardrail_to_apply": _ToolCallBlocker(), "model": "gpt-4"},
+ )
+ async def consume_checked_stream() -> None:
+ async for chunk in stream:
+ assert isinstance(chunk, ModelResponseStream)
+ assert all(not choice.delta.tool_calls and choice.finish_reason is None for choice in chunk.choices)
+
with pytest.raises(GuardrailRaisedException):
- await _drive_stream(UnifiedLLMGuardrails(), _ToolCallBlocker(), [tool_chunk])
+ await consume_checked_stream()
@pytest.mark.asyncio
async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):
diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py
index 508736fb78e..fa7aafa0a72 100644
--- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py
+++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py
@@ -1487,7 +1487,7 @@ async def test_apply_guardrail_invokes_logging_pipeline(mocker):
}
-def _patch_apply_guardrail_env(mocker, guardrail_result):
+def _patch_apply_guardrail_env(mocker, guardrail_result, processed_data=None):
mock_guardrail = mocker.Mock()
mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result)
@@ -1500,7 +1500,7 @@ def _patch_apply_guardrail_env(mocker, guardrail_result):
mock_logging_obj.model_call_details = {}
mock_processor = mocker.Mock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
- return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
+ return_value=(processed_data or {"guardrail_name": "test-guardrail"}, mock_logging_obj)
)
mocker.patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
@@ -1571,6 +1571,68 @@ async def test_apply_guardrail_forwards_metadata_and_messages_together(mocker):
)
+@pytest.mark.asyncio
+async def test_apply_guardrail_replaces_caller_identity_with_the_authenticated_key(mocker):
+ """Identity fields come from the proxy's own sanitized metadata, never from
+ the caller, so a request cannot impersonate another key or probe its policy."""
+ mock_guardrail = _patch_apply_guardrail_env(
+ mocker,
+ {"texts": ["ok"]},
+ processed_data={
+ "guardrail_name": "test-guardrail",
+ "metadata": {
+ "route": "/apply_guardrail",
+ "user_api_key_alias": "billing-app",
+ "user_api_key_user_id": "u-1",
+ },
+ },
+ )
+
+ request = ApplyGuardrailRequest(
+ guardrail_name="test-guardrail",
+ text="hello",
+ metadata={
+ "forbidden_topics": ["tax"],
+ "user_api_key_alias": "someone-else",
+ "user_api_key_user_email": "victim@example.com",
+ },
+ )
+ await apply_guardrail(
+ fastapi_request=mocker.Mock(),
+ request=request,
+ user_api_key_dict=UserAPIKeyAuth(key_alias="billing-app", user_id="u-1"),
+ )
+
+ forwarded = mock_guardrail.apply_guardrail.await_args.kwargs["request_data"]["metadata"]
+ assert forwarded == {
+ "forbidden_topics": ["tax"],
+ "user_api_key_alias": "billing-app",
+ "user_api_key_user_id": "u-1",
+ }
+
+
+@pytest.mark.asyncio
+async def test_apply_guardrail_attaches_key_identity_when_caller_sends_no_metadata(mocker):
+ """A caller that sends no metadata still gets the authenticated identity
+ forwarded, the same shape every LLM route gives a guardrail."""
+ mock_guardrail = _patch_apply_guardrail_env(
+ mocker,
+ {"texts": ["ok"]},
+ processed_data={"guardrail_name": "test-guardrail", "metadata": {"user_api_key_alias": "billing-app"}},
+ )
+
+ request = ApplyGuardrailRequest(guardrail_name="test-guardrail", text="hello")
+ await apply_guardrail(
+ fastapi_request=mocker.Mock(),
+ request=request,
+ user_api_key_dict=UserAPIKeyAuth(key_alias="billing-app"),
+ )
+
+ assert mock_guardrail.apply_guardrail.await_args.kwargs["request_data"] == {
+ "metadata": {"user_api_key_alias": "billing-app"}
+ }
+
+
@pytest.mark.asyncio
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
"""Without metadata, request_data stays empty (backward-compatible)."""
diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py
index d3f1183cea6..a0387f18db2 100644
--- a/tests/unit/interactions/test_openapi_compliance.py
+++ b/tests/unit/interactions/test_openapi_compliance.py
@@ -9,11 +9,19 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
import json
import os
-from typing import Any, Dict
+import re
+from collections.abc import Mapping
+from typing import Any, Dict, Final
from unittest.mock import MagicMock, patch
import httpx
import pytest
+from jsonschema import Draft202012Validator
+from pydantic import TypeAdapter
+
+from litellm.llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig
+from litellm.types.interactions import InteractionInput
+from litellm.types.router import GenericLiteLLMParams
from openapi_core import OpenAPI
OPENAPI_SPEC_URL = "https://ai.google.dev/static/api/interactions.openapi.json"
@@ -56,61 +64,58 @@ def openapi_spec(spec_dict: Dict[str, Any]) -> OpenAPI:
return OpenAPI.from_dict(spec_dict)
+@pytest.fixture(scope="module")
+def model_request_schema(spec_dict: Mapping[str, object]) -> Mapping[str, object]:
+ objects: Final = TypeAdapter(Mapping[str, Mapping[str, object]])
+ components: Final = TypeAdapter(Mapping[str, object]).validate_python(spec_dict["components"])
+ schemas: Final = objects.validate_python(components["schemas"])
+ paths: Final = objects.validate_python(spec_dict["paths"])
+ operation: Final = next(methods["post"] for path, methods in paths.items() if path.endswith("/interactions"))
+ post: Final = TypeAdapter(Mapping[str, object]).validate_python(operation)
+ body: Final = TypeAdapter(Mapping[str, object]).validate_python(post["requestBody"])
+ content: Final = objects.validate_python(body["content"])
+ schema: Final = TypeAdapter(Mapping[str, object]).validate_python(content["application/json"]["schema"])
+ variants: Final = TypeAdapter(tuple[Mapping[str, str], ...]).validate_python(schema["oneOf"])
+ return next(
+ schemas[variant["$ref"].rsplit("/", 1)[-1]]
+ for variant in variants
+ if "model" in TypeAdapter(Mapping[str, object]).validate_python(
+ schemas[variant["$ref"].rsplit("/", 1)[-1]]["properties"]
+ )
+ )
+
+
class TestRequestCompliance:
- """Tests that our request bodies match the OpenAPI spec."""
+ def test_create_model_interaction_request_schema(self, model_request_schema: Mapping[str, object]) -> None:
+ config: Final = GoogleAIStudioInteractionsConfig()
+ payload: Final = config.transform_request(
+ model="test-model", agent=None, input="test input",
+ optional_params={"stream": True, "response_mime_type": "application/json"},
+ litellm_params=GenericLiteLLMParams(api_key="test-key"), headers={},
+ )
+ properties: Final = TypeAdapter(Mapping[str, object]).validate_python(model_request_schema["properties"])
+ assert payload == {
+ "model": "test-model", "input": "test input", "stream": True,
+ "response_format": {"type": "text", "mime_type": "application/json"},
+ }
+ assert payload.keys() <= properties.keys()
+ assert set(config.get_supported_params("test-model")) - {"agent", "response_mime_type"} <= properties.keys()
- def test_create_model_interaction_request_schema(self, spec_dict):
- """Verify CreateModelInteractionParams schema fields."""
- schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
-
- # Required fields per spec
- assert "model" in schema["required"]
- assert "input" in schema["required"]
-
- # Check our supported optional fields exist in spec
- our_optional_fields = [
- "tools",
- "system_instruction",
- "generation_config",
- "stream",
- "store",
- "background",
- "response_modalities",
- "response_format",
- "response_mime_type",
- "previous_interaction_id",
- ]
-
- spec_properties = schema["properties"]
- for field in our_optional_fields:
- assert field in spec_properties, f"Field '{field}' not in OpenAPI spec"
- print(f"✓ Field '{field}' exists in spec")
-
- def test_input_types_match_spec(self, spec_dict):
- """Verify input field supports string, Content, Content[], Turn[]."""
- schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
- input_schema = schema["properties"]["input"]
-
- # The input property may be inline oneOf or a $ref to InteractionsInput
- if "$ref" in input_schema:
- ref_name = input_schema["$ref"].split("/")[-1]
- input_schema = spec_dict["components"]["schemas"][ref_name]
-
- # Should be oneOf with multiple types
- assert "oneOf" in input_schema
-
- input_types = []
- for option in input_schema["oneOf"]:
- if option.get("type") == "string":
- input_types.append("string")
- elif option.get("type") == "array":
- input_types.append("array")
- elif "$ref" in option:
- input_types.append(option["$ref"])
-
- print(f"Input supports types: {input_types}")
- assert "string" in input_types, "Input should support string"
- assert "array" in input_types, "Input should support array"
+ @pytest.mark.parametrize("input_value", ["test input", [{"type": "text", "text": "test input"}]])
+ def test_input_types_match_spec(
+ self, spec_dict: Mapping[str, object], model_request_schema: Mapping[str, object],
+ input_value: InteractionInput,
+ ) -> None:
+ payload: Final = GoogleAIStudioInteractionsConfig().transform_request(
+ model="test-model", agent=None, input=input_value, optional_params={},
+ litellm_params=GenericLiteLLMParams(api_key="test-key"), headers={},
+ )
+ properties: Final = TypeAdapter(Mapping[str, Mapping[str, object]]).validate_python(
+ model_request_schema["properties"]
+ )
+ validator: Final = Draft202012Validator(spec_dict).evolve(schema=properties["input"])
+ assert not tuple(validator.iter_errors(payload["input"]))
+ assert payload["input"] == input_value
def test_content_variants_are_identified_by_their_type_field(self, spec_dict):
"""Verify a Content part can be told apart by its `type`, however the spec spells that.
@@ -307,31 +312,24 @@ class TestEndpointCompliance:
assert create_path is not None, "POST /interactions endpoint not found"
print(f"✓ Create endpoint: POST {create_path}")
- def test_get_endpoint_exists(self, spec_dict):
- """Verify GET /interactions/{id} endpoint exists."""
- paths = spec_dict["paths"]
-
- get_path = None
- for path, methods in paths.items():
- if "{id}" in path and "interactions" in path and "get" in methods:
- get_path = path
- break
-
- assert get_path is not None, "GET /interactions/{id} endpoint not found"
- print(f"✓ Get endpoint: GET {get_path}")
-
- def test_delete_endpoint_exists(self, spec_dict):
- """Verify DELETE /interactions/{id} endpoint exists."""
- paths = spec_dict["paths"]
-
- delete_path = None
- for path, methods in paths.items():
- if "{id}" in path and "interactions" in path and "delete" in methods:
- delete_path = path
- break
-
- assert delete_path is not None, "DELETE /interactions/{id} endpoint not found"
- print(f"✓ Delete endpoint: DELETE {delete_path}")
+ @pytest.mark.parametrize("method", ["get", "delete"])
+ def test_interaction_item_url_matches_spec(self, spec_dict: Mapping[str, object], method: str) -> None:
+ config: Final = GoogleAIStudioInteractionsConfig()
+ transform: Final = (
+ config.transform_get_interaction_request if method == "get" else config.transform_delete_interaction_request
+ )
+ url, body = transform(
+ interaction_id="test-interaction", api_base="https://example.com",
+ litellm_params=GenericLiteLLMParams(api_key="test-key"), headers={},
+ )
+ path: Final = httpx.URL(url).path
+ paths: Final = TypeAdapter(Mapping[str, Mapping[str, object]]).validate_python(spec_dict["paths"])
+ assert any(
+ re.fullmatch(re.sub(r"\{[^}]+\}", "[^/]+", template), path) and method in operations
+ for template, operations in paths.items()
+ ), path
+ assert path.endswith("/test-interaction")
+ assert body == {}
if __name__ == "__main__":
diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py
index 5c85faa5e13..6ba912f25df 100644
--- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py
+++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py
@@ -7,7 +7,7 @@ with guardrail transformations, including tool calls.
import json
from collections.abc import Mapping
-from typing import Any, Literal, Optional
+from typing import Any, Final, Literal, Optional
import pytest
@@ -1395,29 +1395,65 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
)
assert [
- (tool_call["id"], tool_call["function"]["arguments"])
- for tool_call in guardrail.seen_inputs[-1]["tool_calls"]
- ] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')]
+ [(tool_call["id"], tool_call["function"]["arguments"]) for tool_call in inputs["tool_calls"]]
+ for inputs in guardrail.seen_inputs
+ ] == [[("call_1", '{"fruit": "persimmon"}')], [("call_2", '{"fruit": "durian"}')]]
@pytest.mark.asyncio
- async def test_deliver_ended_stream_tool_call_rewrite_on_multi_choice_stream_fails_closed(self):
- from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
+ @pytest.mark.parametrize("transform", [False, True])
+ async def test_each_choice_refreshes_shared_guardrail_response_context(self, transform: bool) -> None:
+ from fastapi import HTTPException
+ from litellm.llms.base_llm.guardrail_translation.base_translation import StreamTransformSink
+ from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
- handler = OpenAIChatCompletionsHandler()
- chunks = self._two_choice_tool_call_stream_chunks()
+ class ContextGuardrail(CustomGuardrail):
+ async def apply_guardrail(
+ self, inputs: GenericGuardrailAPIInputs, request_data: dict[str, object],
+ input_type: Literal["request", "response"], logging_obj: object = None,
+ ) -> GenericGuardrailAPIInputs:
+ response: Final = request_data["response"]
+ assert isinstance(response, ModelResponse)
+ request_data["visited_choices"] = (*request_data.get("visited_choices", ()), response.choices[0].index)
+ if response.choices[0].message.content == "forbidden":
+ raise HTTPException(status_code=400, detail="later choice blocked")
+ return inputs
- with pytest.raises(UndeliverableStreamRewrite, match="the stream carries 2 choices") as raised:
- await handler.process_output_streaming_response(
- responses_so_far=chunks,
- guardrail_to_apply=MockGuardrail(guardrail_name="test"),
- litellm_logging_obj=None,
- deliver_ended_stream_rewrites=True,
+ chunks: Final = [ModelResponseStream(choices=[StreamingChoices(
+ index=index, delta=Delta(content=text, tool_calls=[{
+ "index": 0, "id": f"call_{index}", "type": "function",
+ "function": {"name": "contact", "arguments": "{}"},
+ }]), finish_reason="tool_calls",
+ ) for index, text in enumerate(("allowed", "forbidden"))])]
+ request_data: Final = {"metadata": {"trace": "retained"}}
+ with pytest.raises(HTTPException, match="later choice blocked"):
+ await OpenAIChatCompletionsHandler().process_output_streaming_response(
+ responses_so_far=chunks, guardrail_to_apply=ContextGuardrail(guardrail_name="context"),
+ request_data=request_data, deliver_ended_stream_rewrites=True,
+ stream_transform_sink=StreamTransformSink() if transform else None,
)
+ assert request_data["visited_choices"] == (0, 1)
+ assert request_data["metadata"]["trace"] == "retained"
+ assert request_data["responses"] is chunks
+ restored: Final = request_data["response"]
+ assert isinstance(restored, ModelResponse)
+ assert tuple(choice.index for choice in restored.choices) == (0, 1)
- assert raised.value.guardrail_name == "test"
- assert raised.value.reason == (
- "the stream carries 2 choices and tool-call rewrites are only written back on single-choice streams"
+ @pytest.mark.asyncio
+ async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None:
+ handler: Final = OpenAIChatCompletionsHandler()
+ chunks: Final = self._two_choice_tool_call_stream_chunks()
+ await handler.process_output_streaming_response(
+ responses_so_far=chunks,
+ guardrail_to_apply=MockGuardrail(guardrail_name="test"),
+ litellm_logging_obj=None,
+ deliver_ended_stream_rewrites=True,
)
+ arguments: Final = tuple(
+ "".join(call.function.arguments for chunk in chunks for choice in chunk.choices
+ if choice.index == index for call in choice.delta.tool_calls or ())
+ for index in range(2)
+ )
+ assert arguments == ('{"fruit": "PERSIMMON"}', '{"fruit": "DURIAN"}')
@pytest.mark.asyncio
async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):
diff --git a/tests/unit/rust_bridge/native_route_wheel_test.py b/tests/unit/rust_bridge/native_route_wheel_test.py
index 0b442f1f269..bde7a091a5c 100644
--- a/tests/unit/rust_bridge/native_route_wheel_test.py
+++ b/tests/unit/rust_bridge/native_route_wheel_test.py
@@ -154,12 +154,8 @@ def success_value(route: str, response: dict[object, object]) -> object:
def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
- if route == "chat_completions":
- upstream_error: Final = native.RustUpstreamError
- if not isinstance(error, upstream_error) or error.args[0] != 429:
- raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
- return
- if not isinstance(error, RuntimeError) or "429" not in str(error):
+ upstream_error: Final = native.RustUpstreamError
+ if not isinstance(error, upstream_error) or error.args != (429, '{"error":"native-rate-limit"}'):
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
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 10ca58294b9..bb28d6bc124 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 eb5d47d7891..8fd39120cec 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",
@@ -54,4 +55,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 d88a333d6f1..9c0c2efe6fd 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 476bcd3a8ae..c68b11adfad 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/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx
index 33d35b9677c..b0c8e23a0cb 100644
--- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx
+++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.integration.test.tsx
@@ -157,7 +157,7 @@ describe("EditAutoRouterModal keyword matching", () => {
const threshold = screen.getByRole("textbox", { name: "Success threshold" });
expect(threshold).toHaveValue("0.91");
fireEvent.change(threshold, { target: { value: raw } });
- await waitFor(() => expect(screen.getByRole("button", { name: "Save Changes" })).toBeEnabled());
+ await waitFor(() => expect(screen.getByRole("button", { name: "Save Changes" })).toBeEnabled(), { timeout: 5000 });
await user.click(screen.getByRole("button", { name: "Save Changes" }));
await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalledOnce());
if (raw === "") expect(savedConfig()).not.toHaveProperty("heuristic_v2_success_threshold");
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index 2b8ed9aa58d..75bed9075b8 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -25558,6 +25558,11 @@ export interface components {
* @default true
*/
sticky_session_routing: boolean | null;
+ /**
+ * Streaming Transform Mode
+ * @description Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.
+ */
+ streaming_transform_mode?: ("block_only" | "incremental_diff") | null;
/**
* Template Id
* @description The ID of your Model Armor template
@@ -25570,7 +25575,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', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. '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', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', 'typesafe', 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}
*/
@@ -34245,6 +34250,11 @@ export interface components {
* @description Client secret of the gateway's Entra app registration, used to perform the On-Behalf-Of exchange. Falls back to the AGENT365_CLIENT_SECRET environment variable.
*/
client_secret?: string | null;
+ /**
+ * 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.
@@ -34785,6 +34795,11 @@ export interface components {
* @default true
*/
sticky_session_routing: boolean | null;
+ /**
+ * Streaming Transform Mode
+ * @description Whether a guardrail's text rewrite reaches a streaming client. Implemented by guardrail='prompt_security' and 'neuraltrust'; generic_guardrail_api takes the same setting under optional_params. 'block_only' (default) streams the raw model chunks, so a block still ends the stream but rewrites are dropped. 'incremental_diff' withholds those chunks and streams the guardrail's rewritten text instead. OpenAI chat completions streaming only.
+ */
+ streaming_transform_mode?: ("block_only" | "incremental_diff") | null;
/**
* Template Id
* @description The ID of your Model Armor template