This commit is contained in:
albertbausili 2026-09-28 13:43:22 +02:00 • committed by GitHub
commit 8f408f99ab
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
24 changed files with 2577 additions and 208 deletions

View file

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

View file

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

View file

@ -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": [
{

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

File diff suppressed because it is too large Load diff

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -0,0 +1,22 @@
<svg xmlns="http://www.w3.org/2000/svg" width="32" height="32" viewBox="0 0 32 32" fill="none">
<g clip-path="url(#neuraltrustClip)">
<path fill="url(#neuraltrustGrad)" d="M32 0H0v32h32z" />
<path
fill="#fff"
d="M18.092 20.06a.67.67 0 0 1-.55.3.7.7 0 0 1-.565-.286l-2.704-3.814-1.45 2.103 2.197 3.098a3.08 3.08 0 0 0 2.51 1.297h.038a3.06 3.06 0 0 0 2.502-1.342l8.02-11.477h-2.926z"
/>
<path
fill="#fff"
d="M14.292 11.518a.63.63 0 0 1 .552.286l2.652 3.74 1.449-2.103-2.145-3.024a3.08 3.08 0 0 0-2.509-1.297h-.039a3.06 3.06 0 0 0-2.506 1.35L3.925 21.85l-.085.123h2.91l6.98-10.155a.68.68 0 0 1 .562-.3"
/>
</g>
<defs>
<linearGradient id="neuraltrustGrad" x1="30.667" x2="6.667" y1="0" y2="32" gradientUnits="userSpaceOnUse">
<stop stop-color="#03AFFF" />
<stop offset="1" stop-color="#9B29FF" />
</linearGradient>
<clipPath id="neuraltrustClip">
<path fill="#fff" d="M0 0h32v32H0z" />
</clipPath>
</defs>
</svg>

After

Width:  |  Height:  |  Size: 995 B

View file

@ -216,6 +216,12 @@ export const GUARDRAIL_PRESETS: Record<string, GuardrailPreset> = {
mode: "pre_call",
defaultOn: false,
},
neuraltrust: {
provider: "Neuraltrust",
guardrailNameSuggestion: "NeuralTrust TrustGuard",
mode: "pre_call",
defaultOn: false,
},
noma: {
provider: "Noma",
guardrailNameSuggestion: "Noma Security",

View file

@ -12,6 +12,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record<string, string> = {
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();
});
});

View file

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

View file

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

View file

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

View file

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

View file

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