mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge 7f358aae42 into 2c9b0e00ac
This commit is contained in:
commit
8f408f99ab
24 changed files with 2577 additions and 208 deletions
8
.github/workflows/test-rust.yml
vendored
8
.github/workflows/test-rust.yml
vendored
|
|
@ -88,7 +88,7 @@ jobs:
|
||||||
|
|
||||||
rust-test:
|
rust-test:
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
timeout-minutes: 20
|
timeout-minutes: 30
|
||||||
defaults:
|
defaults:
|
||||||
run:
|
run:
|
||||||
working-directory: litellm-rust
|
working-directory: litellm-rust
|
||||||
|
|
@ -169,7 +169,11 @@ jobs:
|
||||||
env:
|
env:
|
||||||
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
|
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
|
- name: Run pytest tests/test_litellm_rust with the compiled extension
|
||||||
run: make test-rust-extension
|
run: make test-rust-extension
|
||||||
|
|
|
||||||
|
|
@ -18,9 +18,11 @@ import json
|
||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Mapping, Sequence
|
from collections.abc import Mapping, Sequence
|
||||||
|
from itertools import chain
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
from typing import TYPE_CHECKING, Any, Final, Union, cast
|
||||||
|
|
||||||
|
from pydantic import TypeAdapter
|
||||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||||
|
|
||||||
import litellm
|
import litellm
|
||||||
|
|
@ -46,7 +48,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||||
unappliable_request_rewrite,
|
unappliable_request_rewrite,
|
||||||
)
|
)
|
||||||
from litellm.main import stream_chunk_builder
|
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 (
|
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||||
coerce_stream_holdback_value,
|
coerce_stream_holdback_value,
|
||||||
)
|
)
|
||||||
|
|
@ -556,7 +558,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
||||||
terminate the stream. Text rewrites are not propagated to the client here
|
terminate the stream. Text rewrites are not propagated to the client here
|
||||||
(see ``_process_streaming_transform`` for the incremental_diff path) unless
|
(see ``_process_streaming_transform`` for the incremental_diff path) unless
|
||||||
``deliver_ended_stream_rewrites`` opts the ended-stream branch in."""
|
``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:
|
if has_stream_ended:
|
||||||
await self._process_ended_stream(
|
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)
|
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_texts: Final = self._string_choice_contents(model_response)
|
||||||
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
|
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
|
||||||
await self.process_output_response(
|
inspection_responses: Final = (
|
||||||
response=model_response,
|
tuple(
|
||||||
guardrail_to_apply=guardrail_to_apply,
|
model_response.model_copy(update=MappingProxyType({"choices": [choice]}))
|
||||||
litellm_logging_obj=litellm_logging_obj,
|
for choice in model_response.choices
|
||||||
user_api_key_dict=user_api_key_dict,
|
)
|
||||||
request_data=request_data,
|
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:
|
if not deliver_ended_stream_rewrites:
|
||||||
return
|
return
|
||||||
await self._write_ended_stream_text_rewrites(
|
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)
|
self.merge_user_api_key_metadata_into_request(request_data, user_api_key_dict)
|
||||||
|
|
||||||
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
|
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):
|
if responses_so_far and getattr(responses_so_far[0], "model", None):
|
||||||
inputs["model"] = responses_so_far[0].model
|
inputs["model"] = responses_so_far[0].model
|
||||||
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
|
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
|
||||||
|
|
@ -1132,20 +1201,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
||||||
def _function_tool_call_fragments(
|
def _function_tool_call_fragments(
|
||||||
responses_so_far: Sequence["ModelResponseStream"],
|
responses_so_far: Sequence["ModelResponseStream"],
|
||||||
) -> tuple[tuple[ChatCompletionDeltaToolCall, ...], ...]:
|
) -> 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(
|
fragments: Final = tuple(
|
||||||
tool_call
|
(choice.index, tool_call)
|
||||||
for response in responses_so_far
|
for response in responses_so_far
|
||||||
for choice in response.choices
|
for choice in response.choices
|
||||||
for tool_call in choice.delta.tool_calls or ()
|
for tool_call in choice.delta.tool_calls or ()
|
||||||
if isinstance(tool_call, ChatCompletionDeltaToolCall)
|
if isinstance(tool_call, ChatCompletionDeltaToolCall)
|
||||||
)
|
)
|
||||||
identified: Final = frozenset(fragment.index for fragment in fragments if fragment.id)
|
identified: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.id)
|
||||||
named: Final = frozenset(fragment.index for fragment in fragments if fragment.function.name)
|
named: Final = frozenset((choice, fragment.index) for choice, fragment in fragments if fragment.function.name)
|
||||||
return tuple(
|
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(
|
def _write_ended_stream_tool_call_rewrites(
|
||||||
|
|
@ -1155,28 +1222,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
||||||
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
|
pre_guardrail_tool_calls: tuple[tuple[str | None, str], ...],
|
||||||
guardrail_name: str,
|
guardrail_name: str,
|
||||||
) -> None:
|
) -> 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)
|
post_guardrail_tool_calls: Final = self._function_tool_call_shapes(guardrailed_response)
|
||||||
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
|
if post_guardrail_tool_calls == pre_guardrail_tool_calls:
|
||||||
return
|
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)
|
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):
|
if len(fragments_by_tool_call) != len(post_guardrail_tool_calls):
|
||||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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.",
|
"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"
|
"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": {
|
"template_id": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
{
|
{
|
||||||
|
|
@ -10406,7 +10422,7 @@
|
||||||
},
|
},
|
||||||
"unreachable_fallback": {
|
"unreachable_fallback": {
|
||||||
"default": "fail_closed",
|
"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": [
|
"enum": [
|
||||||
"fail_closed",
|
"fail_closed",
|
||||||
"fail_open"
|
"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.",
|
"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"
|
"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": {
|
"confidence_threshold": {
|
||||||
"default": 0.5,
|
"default": 0.5,
|
||||||
"default_value": 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.",
|
"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"
|
"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": {
|
"template_id": {
|
||||||
"anyOf": [
|
"anyOf": [
|
||||||
{
|
{
|
||||||
|
|
|
||||||
|
|
@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, Union,
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
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._logging import verbose_proxy_logger
|
||||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
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:
|
def _execution_timeout_response(timeout: float) -> TestCustomCodeGuardrailResponse:
|
||||||
return TestCustomCodeGuardrailResponse(
|
return TestCustomCodeGuardrailResponse(
|
||||||
success=False,
|
success=False,
|
||||||
|
|
@ -2404,9 +2424,10 @@ async def apply_guardrail(
|
||||||
if litellm_logging_obj is not None:
|
if litellm_logging_obj is not None:
|
||||||
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
|
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
|
||||||
|
|
||||||
|
metadata: Final = _guardrail_request_metadata(request.metadata, data.get("metadata"))
|
||||||
request_data: Final[dict] = {
|
request_data: Final[dict] = {
|
||||||
**({"messages": request.messages} if request.messages is not None else {}),
|
**({"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)
|
_input_type: Final = _resolve_guardrail_input_type(active_guardrail, request.input_type)
|
||||||
guardrailed_inputs: Final = await active_guardrail.apply_guardrail(
|
guardrailed_inputs: Final = await active_guardrail.apply_guardrail(
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
@ -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,
|
||||||
|
}
|
||||||
|
|
@ -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
|
||||||
|
|
@ -708,40 +708,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async for item in response:
|
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):
|
if self._chunk_has_tool_calls(item):
|
||||||
saw_tool_calls = True
|
saw_tool_calls = True
|
||||||
responses_so_far.append(item)
|
responses_so_far.append(item)
|
||||||
last_chunk = 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
|
continue
|
||||||
|
|
||||||
if self._is_trailing_metadata_chunk(item):
|
if self._is_trailing_metadata_chunk(item):
|
||||||
|
|
@ -759,38 +729,62 @@ class UnifiedLLMGuardrails(CustomLogger):
|
||||||
# sampled round here would guardrail the same content twice.
|
# sampled round here would guardrail the same content twice.
|
||||||
if (
|
if (
|
||||||
not end_of_stream_only
|
not end_of_stream_only
|
||||||
|
and not saw_tool_calls
|
||||||
and not self._chunk_has_finish_reason(item)
|
and not self._chunk_has_finish_reason(item)
|
||||||
and chunk_counter % sampling_rate == 0
|
and chunk_counter % sampling_rate == 0
|
||||||
):
|
):
|
||||||
async for out in _round(item, is_final=False):
|
async for out in _round(item, is_final=False):
|
||||||
yield out
|
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:
|
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,
|
endpoint_translation=endpoint_translation,
|
||||||
guardrail_to_apply=guardrail_to_apply,
|
guardrail_to_apply=guardrail_to_apply,
|
||||||
request_data=request_data,
|
request_data=request_data,
|
||||||
user_api_key_dict=user_api_key_dict,
|
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,
|
responses_yielded=responses_yielded,
|
||||||
):
|
):
|
||||||
yield out
|
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(
|
async for out in self._emit_stream_tail(
|
||||||
last_chunk=last_chunk,
|
last_chunk=last_chunk,
|
||||||
final_round=_round,
|
final_round=_round,
|
||||||
|
|
@ -818,7 +812,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
||||||
responses_yielded.append(trailing)
|
responses_yielded.append(trailing)
|
||||||
yield trailing
|
yield trailing
|
||||||
|
|
||||||
async def _inspect_full_response_for_block(
|
async def _inspect_full_response(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
endpoint_translation: _EndpointTranslation,
|
endpoint_translation: _EndpointTranslation,
|
||||||
|
|
@ -828,16 +822,8 @@ class UnifiedLLMGuardrails(CustomLogger):
|
||||||
responses_so_far: Sequence[object],
|
responses_so_far: Sequence[object],
|
||||||
responses_yielded: Sequence[object],
|
responses_yielded: Sequence[object],
|
||||||
) -> AsyncGenerator[object, None]:
|
) -> 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.integrations.custom_guardrail import ModifyResponseException
|
||||||
|
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await endpoint_translation.process_output_streaming_response(
|
await endpoint_translation.process_output_streaming_response(
|
||||||
|
|
@ -847,7 +833,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
||||||
user_api_key_dict=user_api_key_dict,
|
user_api_key_dict=user_api_key_dict,
|
||||||
request_data=request_data,
|
request_data=request_data,
|
||||||
stream_transform_sink=None,
|
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:
|
except ModifyResponseException as e:
|
||||||
if e.original_response is None:
|
if e.original_response is None:
|
||||||
e.original_response = responses_so_far
|
e.original_response = responses_so_far
|
||||||
|
|
|
||||||
|
|
@ -41,6 +41,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
|
||||||
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
|
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
|
||||||
ContentFilterCategoryConfig,
|
ContentFilterCategoryConfig,
|
||||||
)
|
)
|
||||||
|
from litellm.types.proxy.guardrails.guardrail_hooks.neuraltrust import (
|
||||||
|
NeuralTrustGuardrailConfigModel,
|
||||||
|
)
|
||||||
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
|
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
|
||||||
OvalixGuardrailConfigModel,
|
OvalixGuardrailConfigModel,
|
||||||
)
|
)
|
||||||
|
|
@ -78,7 +81,7 @@ Pydantic object defining how to set guardrails on litellm proxy
|
||||||
guardrails:
|
guardrails:
|
||||||
- guardrail_name: "bedrock-pre-guard"
|
- guardrail_name: "bedrock-pre-guard"
|
||||||
litellm_params:
|
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"
|
mode: "during_call"
|
||||||
guardrailIdentifier: ff6ujrregl1q
|
guardrailIdentifier: ff6ujrregl1q
|
||||||
guardrailVersion: "DRAFT"
|
guardrailVersion: "DRAFT"
|
||||||
|
|
@ -96,6 +99,7 @@ class SupportedGuardrailIntegrations(Enum):
|
||||||
PRESIDIO = "presidio"
|
PRESIDIO = "presidio"
|
||||||
HIDE_SECRETS = "hide-secrets"
|
HIDE_SECRETS = "hide-secrets"
|
||||||
HIDDENLAYER = "hiddenlayer"
|
HIDDENLAYER = "hiddenlayer"
|
||||||
|
NEURALTRUST = "neuraltrust"
|
||||||
AIM = "aim"
|
AIM = "aim"
|
||||||
CATO_NETWORKS = "cato_networks"
|
CATO_NETWORKS = "cato_networks"
|
||||||
PANGEA = "pangea"
|
PANGEA = "pangea"
|
||||||
|
|
@ -1059,11 +1063,22 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
||||||
default="fail_closed",
|
default="fail_closed",
|
||||||
description=(
|
description=(
|
||||||
"Behavior when a guardrail endpoint is unreachable due to network errors. "
|
"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."
|
"'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(
|
extra_headers: list[str] | None = Field(
|
||||||
default=None,
|
default=None,
|
||||||
description=(
|
description=(
|
||||||
|
|
@ -1195,6 +1210,7 @@ class LitellmParams( # pyright: ignore[reportIncompatibleVariableOverride] # o
|
||||||
QualifireGuardrailConfigModel,
|
QualifireGuardrailConfigModel,
|
||||||
BlockCodeExecutionGuardrailConfigModel,
|
BlockCodeExecutionGuardrailConfigModel,
|
||||||
HiddenlayerGuardrailConfigModel,
|
HiddenlayerGuardrailConfigModel,
|
||||||
|
NeuralTrustGuardrailConfigModel,
|
||||||
QostodianNexusConfigModel,
|
QostodianNexusConfigModel,
|
||||||
VigilGuardGuardrailConfigModel,
|
VigilGuardGuardrailConfigModel,
|
||||||
SingulrGuardrailConfigModel,
|
SingulrGuardrailConfigModel,
|
||||||
|
|
|
||||||
|
|
@ -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
|
|
@ -3,6 +3,9 @@
|
||||||
import logging
|
import logging
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from typing import TYPE_CHECKING, Final, Literal
|
from typing import TYPE_CHECKING, Final, Literal
|
||||||
|
from collections.abc import AsyncIterator
|
||||||
|
|
||||||
|
from pydantic import TypeAdapter
|
||||||
|
|
||||||
import pytest
|
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.guardrails import GuardrailEventHooks
|
||||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
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:
|
if TYPE_CHECKING:
|
||||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
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 ""
|
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:
|
class TestStreamingTransform:
|
||||||
"""Streaming text-transformation (incremental_diff) path on the OpenAI chat
|
"""Streaming text-transformation (incremental_diff) path on the OpenAI chat
|
||||||
completions streaming surface."""
|
completions streaming surface."""
|
||||||
|
|
@ -947,6 +989,103 @@ class TestStreamingTransform:
|
||||||
def _use_openai_handler_mapping(self, monkeypatch):
|
def _use_openai_handler_mapping(self, monkeypatch):
|
||||||
_patch_translation_mappings(monkeypatch, {CallTypes.acompletion: OpenAIChatCompletionsHandler})
|
_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
|
@pytest.mark.asyncio
|
||||||
async def test_block_only_drops_text_rewrites(self):
|
async def test_block_only_drops_text_rewrites(self):
|
||||||
"""Default block_only: the guardrail's uppercasing never reaches the
|
"""Default block_only: the guardrail's uppercasing never reaches the
|
||||||
|
|
@ -1367,14 +1506,15 @@ class TestStreamingTransform:
|
||||||
assert out[-2].choices[0].finish_reason == "stop"
|
assert out[-2].choices[0].finish_reason == "stop"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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
|
"""A guardrail that blocks on tool calls must terminate the incremental_diff
|
||||||
stream: tool calls go through the block decision, not bypass it."""
|
stream: tool calls go through the block decision, not bypass it."""
|
||||||
from litellm.exceptions import GuardrailRaisedException
|
from litellm.exceptions import GuardrailRaisedException
|
||||||
|
|
||||||
class _ToolCallBlocker(_StreamingTextGuardrail):
|
class _ToolCallBlocker(_StreamingTextGuardrail):
|
||||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
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(
|
raise GuardrailRaisedException(
|
||||||
guardrail_name="tc-block",
|
guardrail_name="tc-block",
|
||||||
message="blocked tool call",
|
message="blocked tool call",
|
||||||
|
|
@ -1387,7 +1527,7 @@ class TestStreamingTransform:
|
||||||
StreamingChoices(
|
StreamingChoices(
|
||||||
index=0,
|
index=0,
|
||||||
delta=Delta(
|
delta=Delta(
|
||||||
content=None,
|
content=content,
|
||||||
tool_calls=[
|
tool_calls=[
|
||||||
{
|
{
|
||||||
"index": 0,
|
"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):
|
with pytest.raises(GuardrailRaisedException):
|
||||||
await _drive_stream(UnifiedLLMGuardrails(), _ToolCallBlocker(), [tool_chunk])
|
await consume_checked_stream()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):
|
async def test_mixed_content_and_tool_call_chunk_does_not_leak_text(self):
|
||||||
|
|
|
||||||
|
|
@ -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 = mocker.Mock()
|
||||||
mock_guardrail.apply_guardrail = AsyncMock(return_value=guardrail_result)
|
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_logging_obj.model_call_details = {}
|
||||||
mock_processor = mocker.Mock()
|
mock_processor = mocker.Mock()
|
||||||
mock_processor.common_processing_pre_call_logic = AsyncMock(
|
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(
|
mocker.patch(
|
||||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
"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
|
@pytest.mark.asyncio
|
||||||
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
|
async def test_apply_guardrail_omits_metadata_when_not_sent(mocker):
|
||||||
"""Without metadata, request_data stays empty (backward-compatible)."""
|
"""Without metadata, request_data stays empty (backward-compatible)."""
|
||||||
|
|
|
||||||
|
|
@ -9,11 +9,19 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import os
|
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
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
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
|
from openapi_core import OpenAPI
|
||||||
|
|
||||||
OPENAPI_SPEC_URL = "https://ai.google.dev/static/api/interactions.openapi.json"
|
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)
|
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:
|
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):
|
@pytest.mark.parametrize("input_value", ["test input", [{"type": "text", "text": "test input"}]])
|
||||||
"""Verify CreateModelInteractionParams schema fields."""
|
def test_input_types_match_spec(
|
||||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
self, spec_dict: Mapping[str, object], model_request_schema: Mapping[str, object],
|
||||||
|
input_value: InteractionInput,
|
||||||
# Required fields per spec
|
) -> None:
|
||||||
assert "model" in schema["required"]
|
payload: Final = GoogleAIStudioInteractionsConfig().transform_request(
|
||||||
assert "input" in schema["required"]
|
model="test-model", agent=None, input=input_value, optional_params={},
|
||||||
|
litellm_params=GenericLiteLLMParams(api_key="test-key"), headers={},
|
||||||
# Check our supported optional fields exist in spec
|
)
|
||||||
our_optional_fields = [
|
properties: Final = TypeAdapter(Mapping[str, Mapping[str, object]]).validate_python(
|
||||||
"tools",
|
model_request_schema["properties"]
|
||||||
"system_instruction",
|
)
|
||||||
"generation_config",
|
validator: Final = Draft202012Validator(spec_dict).evolve(schema=properties["input"])
|
||||||
"stream",
|
assert not tuple(validator.iter_errors(payload["input"]))
|
||||||
"store",
|
assert payload["input"] == input_value
|
||||||
"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"
|
|
||||||
|
|
||||||
def test_content_variants_are_identified_by_their_type_field(self, spec_dict):
|
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.
|
"""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"
|
assert create_path is not None, "POST /interactions endpoint not found"
|
||||||
print(f"✓ Create endpoint: POST {create_path}")
|
print(f"✓ Create endpoint: POST {create_path}")
|
||||||
|
|
||||||
def test_get_endpoint_exists(self, spec_dict):
|
@pytest.mark.parametrize("method", ["get", "delete"])
|
||||||
"""Verify GET /interactions/{id} endpoint exists."""
|
def test_interaction_item_url_matches_spec(self, spec_dict: Mapping[str, object], method: str) -> None:
|
||||||
paths = spec_dict["paths"]
|
config: Final = GoogleAIStudioInteractionsConfig()
|
||||||
|
transform: Final = (
|
||||||
get_path = None
|
config.transform_get_interaction_request if method == "get" else config.transform_delete_interaction_request
|
||||||
for path, methods in paths.items():
|
)
|
||||||
if "{id}" in path and "interactions" in path and "get" in methods:
|
url, body = transform(
|
||||||
get_path = path
|
interaction_id="test-interaction", api_base="https://example.com",
|
||||||
break
|
litellm_params=GenericLiteLLMParams(api_key="test-key"), headers={},
|
||||||
|
)
|
||||||
assert get_path is not None, "GET /interactions/{id} endpoint not found"
|
path: Final = httpx.URL(url).path
|
||||||
print(f"✓ Get endpoint: GET {get_path}")
|
paths: Final = TypeAdapter(Mapping[str, Mapping[str, object]]).validate_python(spec_dict["paths"])
|
||||||
|
assert any(
|
||||||
def test_delete_endpoint_exists(self, spec_dict):
|
re.fullmatch(re.sub(r"\{[^}]+\}", "[^/]+", template), path) and method in operations
|
||||||
"""Verify DELETE /interactions/{id} endpoint exists."""
|
for template, operations in paths.items()
|
||||||
paths = spec_dict["paths"]
|
), path
|
||||||
|
assert path.endswith("/test-interaction")
|
||||||
delete_path = None
|
assert body == {}
|
||||||
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}")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|
|
||||||
|
|
@ -7,7 +7,7 @@ with guardrail transformations, including tool calls.
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from collections.abc import Mapping
|
from collections.abc import Mapping
|
||||||
from typing import Any, Literal, Optional
|
from typing import Any, Final, Literal, Optional
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
|
@ -1395,29 +1395,65 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
|
||||||
)
|
)
|
||||||
|
|
||||||
assert [
|
assert [
|
||||||
(tool_call["id"], tool_call["function"]["arguments"])
|
[(tool_call["id"], tool_call["function"]["arguments"]) for tool_call in inputs["tool_calls"]]
|
||||||
for tool_call in guardrail.seen_inputs[-1]["tool_calls"]
|
for inputs in guardrail.seen_inputs
|
||||||
] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')]
|
] == [[("call_1", '{"fruit": "persimmon"}')], [("call_2", '{"fruit": "durian"}')]]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_deliver_ended_stream_tool_call_rewrite_on_multi_choice_stream_fails_closed(self):
|
@pytest.mark.parametrize("transform", [False, True])
|
||||||
from litellm.proxy.policy_engine.pipeline_executor import UndeliverableStreamRewrite
|
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()
|
class ContextGuardrail(CustomGuardrail):
|
||||||
chunks = self._two_choice_tool_call_stream_chunks()
|
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:
|
chunks: Final = [ModelResponseStream(choices=[StreamingChoices(
|
||||||
await handler.process_output_streaming_response(
|
index=index, delta=Delta(content=text, tool_calls=[{
|
||||||
responses_so_far=chunks,
|
"index": 0, "id": f"call_{index}", "type": "function",
|
||||||
guardrail_to_apply=MockGuardrail(guardrail_name="test"),
|
"function": {"name": "contact", "arguments": "{}"},
|
||||||
litellm_logging_obj=None,
|
}]), finish_reason="tool_calls",
|
||||||
deliver_ended_stream_rewrites=True,
|
) 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"
|
@pytest.mark.asyncio
|
||||||
assert raised.value.reason == (
|
async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None:
|
||||||
"the stream carries 2 choices and tool-call rewrites are only written back on single-choice streams"
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):
|
async def test_deliver_ended_stream_clean_multi_choice_stream_released_untouched(self):
|
||||||
|
|
|
||||||
|
|
@ -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:
|
def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
|
||||||
if route == "chat_completions":
|
upstream_error: Final = native.RustUpstreamError
|
||||||
upstream_error: Final = native.RustUpstreamError
|
if not isinstance(error, upstream_error) or error.args != (429, '{"error":"native-rate-limit"}'):
|
||||||
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):
|
|
||||||
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
|
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
22
ui/litellm-dashboard/public/assets/logos/neuraltrust.svg
Normal file
22
ui/litellm-dashboard/public/assets/logos/neuraltrust.svg
Normal 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 |
|
|
@ -216,6 +216,12 @@ export const GUARDRAIL_PRESETS: Record<string, GuardrailPreset> = {
|
||||||
mode: "pre_call",
|
mode: "pre_call",
|
||||||
defaultOn: false,
|
defaultOn: false,
|
||||||
},
|
},
|
||||||
|
neuraltrust: {
|
||||||
|
provider: "Neuraltrust",
|
||||||
|
guardrailNameSuggestion: "NeuralTrust TrustGuard",
|
||||||
|
mode: "pre_call",
|
||||||
|
defaultOn: false,
|
||||||
|
},
|
||||||
noma: {
|
noma: {
|
||||||
provider: "Noma",
|
provider: "Noma",
|
||||||
guardrailNameSuggestion: "Noma Security",
|
guardrailNameSuggestion: "Noma Security",
|
||||||
|
|
|
||||||
|
|
@ -12,6 +12,7 @@ const EXPECTED_PARTNER_LOGO_FILES: Record<string, string> = {
|
||||||
panw: "palo_alto_networks.jpeg",
|
panw: "palo_alto_networks.jpeg",
|
||||||
cisco_ai_defense: "cisco.png",
|
cisco_ai_defense: "cisco.png",
|
||||||
noma: "noma_security.png",
|
noma: "noma_security.png",
|
||||||
|
neuraltrust: "neuraltrust.svg",
|
||||||
aporia: "aporia.png",
|
aporia: "aporia.png",
|
||||||
aim: "aim_security.jpeg",
|
aim: "aim_security.jpeg",
|
||||||
cato_networks: "cato_networks.svg",
|
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/");
|
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();
|
||||||
|
});
|
||||||
});
|
});
|
||||||
|
|
|
||||||
|
|
@ -319,6 +319,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [
|
||||||
tags: ["Enterprise", "Security", "Prompt Injection", "PII"],
|
tags: ["Enterprise", "Security", "Prompt Injection", "PII"],
|
||||||
providerKey: "CiscoAiDefense",
|
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",
|
id: "noma",
|
||||||
name: "Noma Security",
|
name: "Noma Security",
|
||||||
|
|
|
||||||
|
|
@ -196,6 +196,20 @@ describe("guardrail_info_helpers", () => {
|
||||||
expect(result.logo).toContain("noma_security.png");
|
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", () => {
|
it("should resolve RepelloAI Argus logo and display name", () => {
|
||||||
populateGuardrailProviders({
|
populateGuardrailProviders({
|
||||||
repelloai: { ui_friendly_name: "RepelloAI Argus" },
|
repelloai: { ui_friendly_name: "RepelloAI Argus" },
|
||||||
|
|
|
||||||
|
|
@ -15,6 +15,7 @@ import lakeraAiLogo from "../../../../../public/assets/logos/lakeraai.jpeg";
|
||||||
import lassoLogo from "../../../../../public/assets/logos/lasso.png";
|
import lassoLogo from "../../../../../public/assets/logos/lasso.png";
|
||||||
import litellmLogo from "../../../../../public/assets/logos/litellm_logo.jpg";
|
import litellmLogo from "../../../../../public/assets/logos/litellm_logo.jpg";
|
||||||
import microsoftAzureLogo from "../../../../../public/assets/logos/microsoft_azure.svg";
|
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 nomaSecurityLogo from "../../../../../public/assets/logos/noma_security.png";
|
||||||
import openaiSmallLogo from "../../../../../public/assets/logos/openai_small.svg";
|
import openaiSmallLogo from "../../../../../public/assets/logos/openai_small.svg";
|
||||||
import paloAltoNetworksLogo from "../../../../../public/assets/logos/palo_alto_networks.jpeg";
|
import paloAltoNetworksLogo from "../../../../../public/assets/logos/palo_alto_networks.jpeg";
|
||||||
|
|
@ -187,6 +188,7 @@ export const guardrailLogoMap = {
|
||||||
"Aporia AI": aporiaLogo.src,
|
"Aporia AI": aporiaLogo.src,
|
||||||
"PANW Prisma AIRS": paloAltoNetworksLogo.src,
|
"PANW Prisma AIRS": paloAltoNetworksLogo.src,
|
||||||
"Cisco AI Defense": ciscoLogo.src,
|
"Cisco AI Defense": ciscoLogo.src,
|
||||||
|
NeuralTrust: neuraltrustLogo.src,
|
||||||
"Noma Security": nomaSecurityLogo.src,
|
"Noma Security": nomaSecurityLogo.src,
|
||||||
"Javelin Guardrails": javelinLogo.src,
|
"Javelin Guardrails": javelinLogo.src,
|
||||||
"Pillar Guardrail": pillarLogo.src,
|
"Pillar Guardrail": pillarLogo.src,
|
||||||
|
|
|
||||||
|
|
@ -157,7 +157,7 @@ describe("EditAutoRouterModal keyword matching", () => {
|
||||||
const threshold = screen.getByRole("textbox", { name: "Success threshold" });
|
const threshold = screen.getByRole("textbox", { name: "Success threshold" });
|
||||||
expect(threshold).toHaveValue("0.91");
|
expect(threshold).toHaveValue("0.91");
|
||||||
fireEvent.change(threshold, { target: { value: raw } });
|
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 user.click(screen.getByRole("button", { name: "Save Changes" }));
|
||||||
await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalledOnce());
|
await waitFor(() => expect(modelPatchUpdateCall).toHaveBeenCalledOnce());
|
||||||
if (raw === "") expect(savedConfig()).not.toHaveProperty("heuristic_v2_success_threshold");
|
if (raw === "") expect(savedConfig()).not.toHaveProperty("heuristic_v2_success_threshold");
|
||||||
|
|
|
||||||
17
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
17
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -25558,6 +25558,11 @@ export interface components {
|
||||||
* @default true
|
* @default true
|
||||||
*/
|
*/
|
||||||
sticky_session_routing: boolean | null;
|
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
|
* Template Id
|
||||||
* @description The ID of your Model Armor template
|
* @description The ID of your Model Armor template
|
||||||
|
|
@ -25570,7 +25575,7 @@ export interface components {
|
||||||
timeout?: number | null;
|
timeout?: number | null;
|
||||||
/**
|
/**
|
||||||
* Unreachable Fallback
|
* 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
|
* @default fail_closed
|
||||||
* @enum {string}
|
* @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.
|
* @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;
|
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
|
* Confidence Threshold
|
||||||
* @description Only block or mask when detection confidence >= this value; below threshold, allow or log_only.
|
* @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
|
* @default true
|
||||||
*/
|
*/
|
||||||
sticky_session_routing: boolean | null;
|
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
|
* Template Id
|
||||||
* @description The ID of your Model Armor template
|
* @description The ID of your Model Armor template
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue