mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(guardrails): per-mode stream_scope with bedrock stream and pass-through fixes (#43801)
* feat(guardrails): run each mode only on streaming, non-streaming, or both
Add stream_scope so a rail can target streaming inference, non-streaming inference, or both per pre, during, and post mode
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(guardrails): honor stream_scope in pipelines and dashboard types
Pipeline steps skipped the stream_scope filter, direct construction ignored mixed-case maps, and schema.d.ts was stale.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(guardrails): skip unmatched stream_scope steps instead of allowing
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(ui): keep stored stream_scope keys for modes not on screen
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(ci): format stream_scope helpers and update fork MCP unit tests
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(mcp): keep hang cancellation tests from timing out during setup
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(guardrails): honor path-defined streaming for stream_scope
Passthrough routes like Gemini streamGenerateContent decide streaming from the URL, so stamp that onto hook data before guardrails run.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(rust): copy AnthropicModelCapabilities instead of cloning
Clippy treats clone-on-Copy as an error, which failed rust-lint on the messages request tests.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(guardrails): trust only server stream classification
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(guardrails): keep streaming marker through deepcopy
scan_raw_request snapshots copy each field, so a plain object() marker would lose identity and skip streaming-only rails.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(cost-map): drop duplicate perceptron-mk1.5 row
Two main cost-map PRs both added the OpenRouter model, so the merge left a second key that CI rejects.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(rust): expect native transcription 429 as RustUpstreamError
HTTP status errors from native routes map through route_error_to_pyerr, so the wheel SIGINT child was dying on an outdated RuntimeError check and never reached the hang probe.
Co-authored-by: Cursor <cursoragent@cursor.com>
* fix(tests): follow Google Interactions OpenAPI without hardcoded names
The live spec dropped CreateModelInteractionParams and renamed the item path to {interactionsId}. Misc CI failed because the compliance tests still looked those names up as literals.
* test(guardrails): reproduce stream_scope bedrock and passthrough field gaps
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): cover invalid stored stream_scope reads
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): classify bedrock stream actions and keep caller is_streaming_request
Co-authored-by: Shivi Jain <mobile.350017@gmail.com>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* refactor(guardrails): make stream_scope_allows public and drop mutable builds
Co-authored-by: Shivi Jain <mobile.350017@gmail.com>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): classify pass-through and Bedrock stream scopes
Co-authored-by: Shivi Jain <mobile.350017@gmail.com>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): harden stream classification and validation
Co-authored-by: Shivi Jain <mobile.350017@gmail.com>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): add stream scope integration audit
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): avoid mutating passthrough custom body
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): run stream scope audit without enterprise license
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): keep stored scope restart cell on one worker
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): scope stream scope audit sink assertions to the rail under test
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): assert logging_only scope absence behind an ordered barrier rail
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): assert one logging_only scan per phase after the barrier rail
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): set request scope on pass-through stream fixtures
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): tolerate invalid YAML stream_scope in v1 guardrails list
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* chore: merge main into litellm_guardrail_stream_scope_fixes
Update pass-through pre-call test callbacks for main's endpoint_type argument
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): add request paths to pass-through fixtures
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): cover websocket pass-through stream scope and LIT-9050 outage spend rows
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* chore(lint): remove unused type discipline suppressions
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(guardrails): script the vertex live upstream in the websocket stream scope test
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): keep stream marker json-serializable and restore pass-through helper names
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(lint): allow required Bedrock action re-export
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): strip the stream marker from pass-through payloads
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(guardrails): drop stream marker by value in scans and snapshots, plain-tuple stream scope state
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(ui): render tag-scoped guardrail modes read-only in the custom code editor
A guardrail whose litellm_params.mode is the tag-scoped dict {tags, default}
crashed the Custom Code editor on open: normalizeMode wrapped the dict into
the mode array and StreamScopeFields rendered it as a React child (error #31,
whole dashboard unmounted). Treat a non-string non-array mode as no editable
modes, show formatGuardrailMode(mode) in a disabled input (read-only, matching
the guardrail info view), and keep mode/stream_scope out of the update payload
for such guardrails.
* fix(ui): resolve merge fallout in guardrails components
Deduplicate toModeArray import after the merge, and move the read-only
guardrail details block back into GuardrailReadOnlyDetails (now rendering
the shared mode/logging-only rows plus the stream-scope detail) so
guardrail_info.tsx stays under the 800-line lint budget. Guardrails UI
suite: 298 passed.
* fix(ui): drop duplicate toModeArray import reintroduced by merge
* chore(pass-through): document the deliberate in-place marker strip as mutable-ok
The clear/update on _parsed_body is load-bearing: rebinding to a fresh
mapping instead breaks 75 pass-through tests because the marker-free body
must propagate through the caller's request dict so downstream guardrail
scans and snapshots never observe the server streaming marker.
* fix(guardrails): move stream_scope after timeout in CustomGuardrail init
Inserting stream_scope before the existing timeout parameter shifted the
positional slot of timeout, so positional callers constructed with their
timeout bound to stream_scope (ValueError) and timeout silently None.
Restores the base parameter order; keyword callers are unaffected.
* fix(guardrails): typing pass for the lint gates
stream_scope leaves the declared constructor parameters (restoring the
base positional surface; keyword construction unchanged), unknown config
values crossing the new stream-scope code paths get typed locals or
cast-ok boundaries, and the passthrough payload literals are annotated.
All three lint gates pass against current main; the guardrail suites are
unchanged (236+5019 passing; one known anyio-driver failure pre-existing
on base).
* style: sort cast imports for the ruff gate
* chore(pass-through): drop unused BEDROCK_STREAMING_ACTIONS re-export
The streaming check now uses is_bedrock_streaming_endpoint; nothing in
the repo imports the name from this module.
---------
Co-authored-by: Shivi Jain <mobile.350017@gmail.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: gabriele <gabriele@berri.ai>
This commit is contained in:
parent
7ec2a94cd8
commit
85a3869dfd
42 changed files with 6315 additions and 367 deletions
|
|
@ -1,10 +1,20 @@
|
|||
import os
|
||||
import sys
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_range, get_env_int_or_none
|
||||
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification"
|
||||
|
||||
|
||||
class ServerStreamingClassification(str, Enum):
|
||||
MARKER = "litellm-server-streaming"
|
||||
|
||||
|
||||
SERVER_STREAMING_CLASSIFICATION_MARKER: Final = ServerStreamingClassification.MARKER
|
||||
|
||||
DEFER_PYDANTIC_BUILD: Final = os.getenv("DEFER_PYDANTIC_BUILD", "true") in ("true", "1", "on")
|
||||
DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"))
|
||||
AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview"))
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import secrets
|
|||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, cast, get_args
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -21,11 +21,14 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
)
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
from litellm.types.guardrails import (
|
||||
DEFAULT_GUARDRAIL_STREAM_SCOPE,
|
||||
DynamicGuardrailParams,
|
||||
GuardrailEventHooks,
|
||||
GuardrailStreamScope,
|
||||
LitellmParams,
|
||||
LoggingOnlyScope,
|
||||
Mode,
|
||||
runtime_stream_scope,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
|
@ -49,6 +52,8 @@ from litellm.constants import (
|
|||
GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS,
|
||||
LOGS_GUARDRAIL_INFORMATION_MARKER,
|
||||
PRE_CALL_EXECUTED_GUARDRAILS_KEY,
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY,
|
||||
SERVER_STREAMING_CLASSIFICATION_MARKER,
|
||||
)
|
||||
from litellm.exceptions import (
|
||||
BlockedPiiEntityError,
|
||||
|
|
@ -173,6 +178,42 @@ def get_session_id_from_request_data(request_data: dict[str, Any]) -> str | None
|
|||
return None
|
||||
|
||||
|
||||
_REALTIME_STREAMING_HOOKS: Final = frozenset({GuardrailEventHooks.realtime_input_transcription})
|
||||
|
||||
|
||||
def without_server_streaming_classification(data: Mapping[str, object]) -> dict[str, object]:
|
||||
return {
|
||||
key: value
|
||||
for key, value in data.items()
|
||||
if key != SERVER_STREAMING_CLASSIFICATION_KEY or value != SERVER_STREAMING_CLASSIFICATION_MARKER
|
||||
}
|
||||
|
||||
|
||||
def guardrail_request_data_with_streaming(
|
||||
data: Mapping[str, object],
|
||||
*,
|
||||
is_streaming: bool,
|
||||
) -> dict[str, object]:
|
||||
data_without_server_classification: Final = without_server_streaming_classification(data)
|
||||
if not is_streaming:
|
||||
return data_without_server_classification
|
||||
return {
|
||||
**data_without_server_classification,
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER,
|
||||
}
|
||||
|
||||
|
||||
def _request_is_streaming(data: object, event_type: GuardrailEventHooks | None = None) -> bool:
|
||||
if event_type in _REALTIME_STREAMING_HOOKS:
|
||||
return True
|
||||
if not isinstance(data, Mapping):
|
||||
return False
|
||||
return (
|
||||
data.get("stream") is True
|
||||
or data.get(SERVER_STREAMING_CLASSIFICATION_KEY) is SERVER_STREAMING_CLASSIFICATION_MARKER
|
||||
)
|
||||
|
||||
|
||||
class CustomGuardrail(CustomLogger):
|
||||
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
|
||||
use_native_during_call_hook: ClassVar[bool] = False
|
||||
|
|
@ -183,6 +224,9 @@ class CustomGuardrail(CustomLogger):
|
|||
records_own_guardrail_information: ClassVar[bool] = False
|
||||
logging_only_scope: LoggingOnlyScope | None
|
||||
|
||||
stream_scope_default: GuardrailStreamScope = DEFAULT_GUARDRAIL_STREAM_SCOPE
|
||||
stream_scope_by_hook: tuple[tuple[str, GuardrailStreamScope], ...] = ()
|
||||
|
||||
timeout: float | httpx.Timeout | None = None
|
||||
|
||||
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
|
||||
|
|
@ -258,6 +302,8 @@ class CustomGuardrail(CustomLogger):
|
|||
self.run_in_parallel: bool = run_in_parallel
|
||||
self.scan_raw_request: bool = scan_raw_request
|
||||
self.only_scan_new_messages: bool = only_scan_new_messages
|
||||
stream_scope_arg: Final[object] = cast(object, kwargs.pop("stream_scope", None)) # cast-ok: config
|
||||
self.apply_stream_scope(stream_scope_arg)
|
||||
self.logging_only_scope = None
|
||||
if timeout is not None:
|
||||
self.timeout = timeout
|
||||
|
|
@ -1099,6 +1145,23 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
return name in suppressed_compression_guardrails()
|
||||
|
||||
def apply_stream_scope(self, stream_scope: object) -> None:
|
||||
default, by_hook = runtime_stream_scope(stream_scope)
|
||||
self.stream_scope_default = default
|
||||
self.stream_scope_by_hook = tuple(by_hook.items())
|
||||
|
||||
def stream_scope_allows(self, data: object, event_type: GuardrailEventHooks) -> bool:
|
||||
scope: Final = next(
|
||||
(scope for hook, scope in self.stream_scope_by_hook if hook == event_type.value),
|
||||
self.stream_scope_default,
|
||||
)
|
||||
if scope == "both":
|
||||
return True
|
||||
is_streaming: Final = _request_is_streaming(data, event_type)
|
||||
if scope == "streaming":
|
||||
return is_streaming
|
||||
return not is_streaming
|
||||
|
||||
def should_run_guardrail(
|
||||
self,
|
||||
data,
|
||||
|
|
@ -1142,8 +1205,10 @@ class CustomGuardrail(CustomLogger):
|
|||
data, self.event_hook, event_type
|
||||
)
|
||||
if result is not None:
|
||||
return result
|
||||
return True
|
||||
tagged_result: Final[bool] = bool(cast(object, result)) # cast-ok: helper return
|
||||
data_obj: Final[object] = cast(object, data) # cast-ok: data param
|
||||
return tagged_result and self.stream_scope_allows(data_obj, event_type)
|
||||
return self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data param
|
||||
return False
|
||||
|
||||
if (
|
||||
|
|
@ -1167,8 +1232,9 @@ class CustomGuardrail(CustomLogger):
|
|||
)
|
||||
result = EnterpriseCustomGuardrailHelper._should_run_if_mode_by_tag(data, self.event_hook, event_type)
|
||||
if result is not None:
|
||||
return result
|
||||
return True
|
||||
mode_tag_result: Final[bool] = bool(cast(object, result)) # cast-ok: helper return
|
||||
return mode_tag_result and self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data
|
||||
return self.stream_scope_allows(cast(object, data), event_type) # cast-ok: data param
|
||||
|
||||
def _event_hook_is_event_type(self, event_type: GuardrailEventHooks) -> bool:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -782,7 +782,7 @@ class RealTimeStreaming:
|
|||
isinstance(cb, CustomGuardrail)
|
||||
and any(
|
||||
cb.should_run_guardrail(
|
||||
data=self.request_data,
|
||||
data={**self.request_data, "stream": True},
|
||||
event_type=et,
|
||||
)
|
||||
for et in event_hooks
|
||||
|
|
@ -847,7 +847,7 @@ class RealTimeStreaming:
|
|||
if event_hooks is None:
|
||||
event_hooks = [GuardrailEventHooks.realtime_input_transcription]
|
||||
_realtime_event_types: Final = event_hooks
|
||||
_check_data: Final = {**self.request_data, "transcript": transcript}
|
||||
_check_data: Final = {**self.request_data, "transcript": transcript, "stream": True}
|
||||
_already_run: Final[set] = set()
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
|
|
|
|||
|
|
@ -22,6 +22,13 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import CostResponseTypes
|
||||
|
||||
|
||||
BEDROCK_STREAMING_ACTIONS: Final = frozenset({"invoke-with-response-stream", "converse-stream"})
|
||||
|
||||
|
||||
def is_bedrock_streaming_endpoint(endpoint: str) -> bool:
|
||||
return endpoint.partition("?")[0].rstrip("/").rsplit("/", 1)[-1] in BEDROCK_STREAMING_ACTIONS
|
||||
|
||||
|
||||
_TEXT_ONLY_DELTA_FIELDS: Final = frozenset({"content", "role"})
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from collections.abc import Mapping
|
|||
from typing import TYPE_CHECKING, Any, Final, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import without_server_streaming_classification
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.proxy._types import PassThroughGuardrailSettings
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
|
@ -80,7 +81,9 @@ class PassThroughEndpointHandler(BaseTranslation):
|
|||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
payload_to_check: Final = {
|
||||
k: v for k, v in data.items() if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj")
|
||||
k: v
|
||||
for k, v in without_server_streaming_classification(data).items()
|
||||
if not k.startswith("_") and k not in ("metadata", "litellm_logging_obj")
|
||||
}
|
||||
verbose_proxy_logger.debug("PassThroughEndpointHandler: Using full payload for guardrail")
|
||||
return safe_dumps(payload_to_check)
|
||||
|
|
|
|||
|
|
@ -11906,6 +11906,34 @@
|
|||
"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"
|
||||
},
|
||||
"stream_scope": {
|
||||
"anyOf": [
|
||||
{
|
||||
"enum": [
|
||||
"streaming",
|
||||
"non_streaming",
|
||||
"both"
|
||||
],
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"additionalProperties": {
|
||||
"enum": [
|
||||
"streaming",
|
||||
"non_streaming",
|
||||
"both"
|
||||
],
|
||||
"type": "string"
|
||||
},
|
||||
"type": "object"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.",
|
||||
"title": "Stream Scope"
|
||||
},
|
||||
"template_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -14831,6 +14859,34 @@
|
|||
"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"
|
||||
},
|
||||
"stream_scope": {
|
||||
"anyOf": [
|
||||
{
|
||||
"enum": [
|
||||
"streaming",
|
||||
"non_streaming",
|
||||
"both"
|
||||
],
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"additionalProperties": {
|
||||
"enum": [
|
||||
"streaming",
|
||||
"non_streaming",
|
||||
"both"
|
||||
],
|
||||
"type": "string"
|
||||
},
|
||||
"type": "object"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.",
|
||||
"title": "Stream Scope"
|
||||
},
|
||||
"template_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ from litellm.types.guardrails import (
|
|||
PresidioPresidioConfigModelUserInterface,
|
||||
SupportedGuardrailIntegrations,
|
||||
ToolPermissionGuardrailConfigModel,
|
||||
with_tolerated_stream_scope,
|
||||
)
|
||||
from litellm.types.llms.base import LiteLLMBaseModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.hide_secrets import (
|
||||
|
|
@ -146,7 +147,7 @@ def _get_guardrails_list_response(
|
|||
GuardrailInfoResponse(
|
||||
guardrail_id=guardrail.get("guardrail_id"),
|
||||
guardrail_name=guardrail.get("guardrail_name"),
|
||||
litellm_params=masked_params,
|
||||
litellm_params=with_tolerated_stream_scope(masked_params),
|
||||
guardrail_info=guardrail.get("guardrail_info"),
|
||||
)
|
||||
)
|
||||
|
|
@ -289,7 +290,7 @@ async def list_guardrails_v2(
|
|||
)
|
||||
masked_litellm_params = (
|
||||
parse_tolerant_litellm_params(
|
||||
masked_litellm_params_dict,
|
||||
with_tolerated_stream_scope(masked_litellm_params_dict),
|
||||
guardrail.get("guardrail_name") or "Unknown",
|
||||
params_model=BaseLitellmParams,
|
||||
)
|
||||
|
|
@ -336,7 +337,7 @@ async def list_guardrails_v2(
|
|||
)
|
||||
masked_in_memory_litellm_params_typed = (
|
||||
parse_tolerant_litellm_params(
|
||||
masked_in_memory_litellm_params,
|
||||
with_tolerated_stream_scope(masked_in_memory_litellm_params),
|
||||
guardrail.get("guardrail_name") or "Unknown",
|
||||
params_model=BaseLitellmParams,
|
||||
)
|
||||
|
|
@ -1263,7 +1264,7 @@ async def patch_guardrail(
|
|||
# Update litellm_params if default_on is provided or pii_entities_config is provided
|
||||
existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {})))
|
||||
current_litellm_params: Final = parse_tolerant_litellm_params(
|
||||
existing_litellm_params,
|
||||
with_tolerated_stream_scope(existing_litellm_params),
|
||||
existing_guardrail.get("guardrail_name") or "Unknown",
|
||||
)
|
||||
requested_litellm_params: Final[Mapping[str, object]] = (
|
||||
|
|
@ -1275,7 +1276,7 @@ async def patch_guardrail(
|
|||
MappingProxyType({**current_litellm_params.model_dump(exclude_unset=True), **requested_litellm_params})
|
||||
)
|
||||
try:
|
||||
parsed_litellm_params: Final = LitellmParams(**merged_litellm_params)
|
||||
parsed_litellm_params: Final = LitellmParams(**with_tolerated_stream_scope(merged_litellm_params))
|
||||
except ValidationError as validation_error:
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
|
|
@ -1436,7 +1437,7 @@ async def get_guardrail_info(guardrail_id: str):
|
|||
)
|
||||
masked_litellm_params = (
|
||||
parse_tolerant_litellm_params(
|
||||
masked_litellm_params_dict,
|
||||
with_tolerated_stream_scope(masked_litellm_params_dict),
|
||||
result.get("guardrail_name") or "Unknown",
|
||||
params_model=BaseLitellmParams,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ from litellm.types.guardrails import (
|
|||
LakeraCategoryThresholds,
|
||||
LitellmParams,
|
||||
SupportedGuardrailIntegrations,
|
||||
with_tolerated_stream_scope,
|
||||
)
|
||||
|
||||
from .guardrail_hooks.llm_as_a_judge import (
|
||||
|
|
@ -634,6 +635,8 @@ def _configure_callback_scoping(
|
|||
"skip_tool_message_in_guardrail are enabled together, which excludes every message from "
|
||||
"scanning, so no request content would ever be scanned. Remove one of the two."
|
||||
)
|
||||
if isinstance(custom_guardrail_callback, CustomGuardrail): # pyright: ignore[reportUnnecessaryIsInstance] # module-path classes may only subclass CustomLogger
|
||||
custom_guardrail_callback.apply_stream_scope(litellm_params.stream_scope)
|
||||
_apply_configured_bool_overrides(custom_guardrail_callback, litellm_params)
|
||||
|
||||
|
||||
|
|
@ -718,9 +721,11 @@ class InMemoryGuardrailHandler:
|
|||
|
||||
if isinstance(litellm_params_data, dict):
|
||||
if reject_invalid_logging_only_scope:
|
||||
litellm_params = LitellmParams(**litellm_params_data)
|
||||
litellm_params = LitellmParams(**with_tolerated_stream_scope(litellm_params_data))
|
||||
else:
|
||||
litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"])
|
||||
litellm_params = parse_tolerant_litellm_params(
|
||||
with_tolerated_stream_scope(litellm_params_data), guardrail["guardrail_name"]
|
||||
)
|
||||
else:
|
||||
litellm_params = litellm_params_data
|
||||
|
||||
|
|
@ -863,14 +868,17 @@ class InMemoryGuardrailHandler:
|
|||
# Extract additional params from litellm_params to pass to custom guardrail
|
||||
# This matches the behavior of other guardrail initializers (e.g., initialize_lakera)
|
||||
# and aligns with the documented behavior for custom guardrails
|
||||
if hasattr(litellm_params, "model_dump"):
|
||||
extra_params = litellm_params.model_dump(exclude_none=True)
|
||||
else:
|
||||
extra_params = dict(litellm_params) if litellm_params else {}
|
||||
|
||||
# Remove params that are handled explicitly or are internal
|
||||
for key in ["guardrail", "mode", "default_on"]:
|
||||
extra_params.pop(key, None)
|
||||
excluded_extra_param_keys: Final = frozenset(("guardrail", "mode", "default_on", "stream_scope"))
|
||||
extra_params_items: Final = (
|
||||
litellm_params.model_dump(exclude_none=True).items()
|
||||
if hasattr(litellm_params, "model_dump")
|
||||
else iter(litellm_params)
|
||||
if litellm_params
|
||||
else ()
|
||||
)
|
||||
extra_params: Final = MappingProxyType(
|
||||
{key: value for key, value in extra_params_items if key not in excluded_extra_param_keys}
|
||||
)
|
||||
|
||||
_guardrail_callback: Final = _guardrail_class(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
|
|
@ -1007,7 +1015,7 @@ class InMemoryGuardrailHandler:
|
|||
return params.model_dump()
|
||||
if isinstance(params, dict):
|
||||
try:
|
||||
return parse_tolerant_litellm_params(params, guardrail_name).model_dump()
|
||||
return parse_tolerant_litellm_params(with_tolerated_stream_scope(params), guardrail_name).model_dump()
|
||||
except ValidationError as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s",
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from litellm.constants import (
|
|||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
X_LITELLM_DISABLE_CALLBACKS,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import without_server_streaming_classification
|
||||
from litellm.litellm_core_utils.core_helpers import is_codex_user_agent
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
|
|
@ -2094,7 +2095,9 @@ def refresh_proxy_server_request_body_snapshot(
|
|||
| _TRANSPORT_ONLY_CREDENTIAL_KEYS
|
||||
| _CALLBACK_CREDENTIAL_KEYS
|
||||
)
|
||||
body: Final = {k: v for k, v in data.items() if k not in _body_snapshot_exclude}
|
||||
body: Final = {
|
||||
k: v for k, v in without_server_streaming_classification(data).items() if k not in _body_snapshot_exclude
|
||||
}
|
||||
proxy_server_request["body"] = body
|
||||
if guardrails_applied and isinstance(logging_obj, Logging):
|
||||
metadata: Final = data.get(get_metadata_variable_name_from_kwargs(data))
|
||||
|
|
|
|||
|
|
@ -44,12 +44,14 @@ from litellm.constants import (
|
|||
AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER,
|
||||
BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import guardrail_request_data_with_streaming
|
||||
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers
|
||||
from litellm.llms.azure.passthrough.transformation import (
|
||||
foreign_azure_deployment,
|
||||
is_azure_body_model_inference_endpoint,
|
||||
)
|
||||
from litellm.llms.bedrock.passthrough.transformation import is_bedrock_streaming_endpoint
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.llms.deepgram.common_utils import (
|
||||
deepgram_listen_callback_params,
|
||||
|
|
@ -946,8 +948,6 @@ BEDROCK_ENDPOINT_ACTIONS: Final = {
|
|||
"count-tokens",
|
||||
}
|
||||
|
||||
BEDROCK_STREAMING_ACTIONS: Final = {"invoke-with-response-stream", "converse-stream"}
|
||||
|
||||
|
||||
def is_bedrock_count_tokens_endpoint(endpoint: str) -> bool:
|
||||
return "count_tokens" in endpoint or "count-tokens" in endpoint
|
||||
|
|
@ -1077,7 +1077,7 @@ async def handle_bedrock_passthrough_router_model(
|
|||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
# Detect streaming based on endpoint
|
||||
is_streaming: Final = any(action in endpoint for action in BEDROCK_STREAMING_ACTIONS)
|
||||
is_streaming: Final = is_bedrock_streaming_endpoint(endpoint)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Bedrock router passthrough: model='%s', endpoint='%s', streaming=%s", model, endpoint, is_streaming
|
||||
|
|
@ -1085,15 +1085,19 @@ async def handle_bedrock_passthrough_router_model(
|
|||
|
||||
# Use the common processing path (same as non-router models)
|
||||
# This ensures all metadata, hooks, and logging are properly initialized
|
||||
data: Final[dict[str, object]] = {}
|
||||
bedrock_payload: Final[dict[str, object]] = {
|
||||
"model": model,
|
||||
"method": request.method,
|
||||
"endpoint": endpoint,
|
||||
"data": request_body,
|
||||
"custom_llm_provider": "bedrock",
|
||||
}
|
||||
data: Final[dict[str, object]] = guardrail_request_data_with_streaming(
|
||||
MappingProxyType(bedrock_payload),
|
||||
is_streaming=is_streaming,
|
||||
)
|
||||
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
|
||||
data["model"] = model
|
||||
data["method"] = request.method
|
||||
data["endpoint"] = endpoint
|
||||
data["data"] = request_body
|
||||
data["custom_llm_provider"] = "bedrock"
|
||||
|
||||
# Use the common passthrough processing to handle metadata and hooks
|
||||
# This also handles all response formatting (streaming/non-streaming) and exceptions
|
||||
try:
|
||||
|
|
@ -1285,14 +1289,19 @@ async def bedrock_llm_proxy_route(
|
|||
"Bedrock passthrough: Using direct Bedrock model '%s' for endpoint '%s'", model, endpoint
|
||||
)
|
||||
|
||||
data: Final[dict[str, object]] = {}
|
||||
is_streaming: Final = is_bedrock_streaming_endpoint(endpoint)
|
||||
passthrough_payload: Final[dict[str, object]] = {
|
||||
"method": request.method,
|
||||
"endpoint": endpoint,
|
||||
"data": request_body,
|
||||
"custom_llm_provider": "bedrock",
|
||||
}
|
||||
data: Final[dict[str, object]] = guardrail_request_data_with_streaming(
|
||||
MappingProxyType(passthrough_payload),
|
||||
is_streaming=is_streaming,
|
||||
)
|
||||
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
|
||||
data["method"] = request.method
|
||||
data["endpoint"] = endpoint
|
||||
data["data"] = request_body
|
||||
data["custom_llm_provider"] = "bedrock"
|
||||
|
||||
try:
|
||||
result: Final = await base_llm_response_processor.base_passthrough_process_llm_request(
|
||||
request=request,
|
||||
|
|
|
|||
|
|
@ -48,7 +48,11 @@ from litellm.constants import (
|
|||
SESSION_ID_OMITTED_METADATA_KEY,
|
||||
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
guardrail_request_data_with_streaming,
|
||||
without_server_streaming_classification,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
bind_budget_reservation_to_callbacks,
|
||||
|
|
@ -603,6 +607,12 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
_parsed_body = _parsed_body or {}
|
||||
parsed_body_typed: Final[Mapping[str, object]] = cast(Mapping[str, object], _parsed_body) # cast-ok: json
|
||||
server_marker_free_body: Final = without_server_streaming_classification(parsed_body_typed)
|
||||
# The marker-free body must propagate through the caller's request dict, so
|
||||
# downstream guardrail scans and snapshots never observe the server streaming marker.
|
||||
_parsed_body.clear()
|
||||
_parsed_body.update(server_marker_free_body)
|
||||
managed_model: Final = get_model_from_request(
|
||||
request_data=_parsed_body,
|
||||
route=get_request_route(request),
|
||||
|
|
@ -828,7 +838,7 @@ def _build_passthrough_failure_request_payload(
|
|||
error response. Spend tracking only attributes a recovered cost when it
|
||||
comes paired with a usage object, so both keys are written together.
|
||||
"""
|
||||
request_payload: Final[dict] = dict(parsed_body or {})
|
||||
request_payload: Final[dict] = dict(cast(Mapping[str, object], parsed_body or {})) # cast-ok: json body
|
||||
if kwargs:
|
||||
request_payload.update(kwargs)
|
||||
if logging_obj is not None:
|
||||
|
|
@ -1163,7 +1173,7 @@ async def pass_through_request(
|
|||
is_multipart: Final = HttpPassThroughEndpointHelpers.is_multipart(request) and not custom_body
|
||||
|
||||
if custom_body:
|
||||
_parsed_body = custom_body
|
||||
_parsed_body = dict(custom_body)
|
||||
elif is_multipart:
|
||||
# Don't parse multipart body here - it will be handled by make_multipart_http_request
|
||||
_parsed_body = {}
|
||||
|
|
@ -1232,6 +1242,14 @@ async def pass_through_request(
|
|||
if _parsed_body is None:
|
||||
_parsed_body = {}
|
||||
_parsed_body["litellm_logging_obj"] = logging_obj
|
||||
is_streaming_pass_through: Final = bool(
|
||||
HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
|
||||
parsed_body=_parsed_body,
|
||||
stream=stream,
|
||||
)
|
||||
)
|
||||
typed_body: Final[Mapping[str, object]] = cast(Mapping[str, object], _parsed_body) # cast-ok: json
|
||||
_parsed_body = guardrail_request_data_with_streaming(typed_body, is_streaming=is_streaming_pass_through)
|
||||
|
||||
### CALL HOOKS ### - modify incoming data / reject request before calling the model
|
||||
_parsed_body = await proxy_logging_obj.pre_call_hook(
|
||||
|
|
@ -2518,10 +2536,10 @@ async def websocket_passthrough_request(
|
|||
)
|
||||
|
||||
### CALL HOOKS ### - modify incoming data / reject request before calling the model
|
||||
websocket_data: dict[str, object] = {}
|
||||
websocket_data = await proxy_logging_obj.pre_call_hook(
|
||||
websocket_hook_data: Final = guardrail_request_data_with_streaming(MappingProxyType({}), is_streaming=True)
|
||||
await proxy_logging_obj.pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=websocket_data,
|
||||
data=websocket_hook_data,
|
||||
call_type="pass_through_endpoint",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,8 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding.
|
|||
import copy
|
||||
import time
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final, Literal, TypeVar
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, TypeVar, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -29,6 +30,7 @@ from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_g
|
|||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import (
|
||||
PipelineExecutionResult,
|
||||
PipelineStep,
|
||||
|
|
@ -265,6 +267,30 @@ class _LegacyHookStreamAdapter(CustomGuardrail):
|
|||
return recorder.inputs
|
||||
|
||||
|
||||
_PIPELINE_EVENT_HOOKS: Final = MappingProxyType(
|
||||
{
|
||||
"pre_call": GuardrailEventHooks.pre_call,
|
||||
"post_call": GuardrailEventHooks.post_call,
|
||||
"during_call": GuardrailEventHooks.during_call,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _pipeline_stream_scope_allows(
|
||||
callback: CustomGuardrail,
|
||||
hook_input: Mapping[str, object],
|
||||
mode: str,
|
||||
streaming_chunks: list[object] | None,
|
||||
) -> bool:
|
||||
event_type: Final = _PIPELINE_EVENT_HOOKS.get(mode)
|
||||
if event_type is None:
|
||||
return True
|
||||
return callback.stream_scope_allows(
|
||||
hook_input if streaming_chunks is None else {**hook_input, "stream": True},
|
||||
event_type,
|
||||
)
|
||||
|
||||
|
||||
def _prepare_hook_input(
|
||||
step: PipelineStep,
|
||||
callback: CustomGuardrail,
|
||||
|
|
@ -494,7 +520,7 @@ class PipelineExecutor:
|
|||
streaming_chunks: list[object] | None = None, # mutable-ok: shared buffered-stream chunks, read per step
|
||||
endpoint_translation: "BaseTranslation | None" = None,
|
||||
) -> tuple[
|
||||
Literal["pass", "fail", "error"],
|
||||
Literal["pass", "fail", "error", "skip"],
|
||||
dict | None,
|
||||
str | None,
|
||||
Exception | None,
|
||||
|
|
@ -504,7 +530,7 @@ class PipelineExecutor:
|
|||
|
||||
Returns:
|
||||
Tuple of (outcome, modified_data, error_detail, original_exception):
|
||||
- outcome: "pass", "fail", or "error"
|
||||
- outcome: "pass", "fail", "error", or "skip"
|
||||
- modified_data: dict if guardrail returned modified data, else None
|
||||
- error_detail: error message string if fail/error, else None
|
||||
- original_exception: the exception the guardrail raised, so the
|
||||
|
|
@ -516,6 +542,10 @@ class PipelineExecutor:
|
|||
verbose_proxy_logger.warning("Pipeline: guardrail '%s' not found in callbacks", step.guardrail)
|
||||
return ("error", None, f"Guardrail '{step.guardrail}' not found", None)
|
||||
|
||||
hook_data: Final[Mapping[str, object]] = cast(Mapping[str, object], data) # cast-ok: payload
|
||||
if not _pipeline_stream_scope_allows(callback, hook_data, mode, streaming_chunks):
|
||||
return ("skip", None, None, None)
|
||||
|
||||
hook_input, scans_raw_request = _prepare_hook_input(step, callback, data, raw_request_snapshot)
|
||||
snapshot_entries_before: Final = len(_recorded_guardrail_information(hook_input))
|
||||
|
||||
|
|
@ -702,10 +732,13 @@ def _pipeline_action_for_outcome(step: PipelineStep, outcome: str) -> str:
|
|||
"""
|
||||
Map pipeline step outcome to the configured action.
|
||||
|
||||
- skip -> next (stream_scope mismatch; do not apply on_pass/on_fail)
|
||||
- pass -> on_pass
|
||||
- fail -> on_fail (content/policy intervention)
|
||||
- error -> on_error if set, else on_fail (backward compatible)
|
||||
"""
|
||||
if outcome == "skip":
|
||||
return "next"
|
||||
if outcome == "pass":
|
||||
return step.on_pass
|
||||
if outcome == "fail":
|
||||
|
|
|
|||
|
|
@ -2,11 +2,12 @@ from collections.abc import Mapping
|
|||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from typing import Final, Literal, cast
|
||||
|
||||
from pydantic import ConfigDict, Field, field_validator, model_validator
|
||||
from typing_extensions import ReadOnly, Required, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
|
||||
from litellm.types.llms.base import LiteLLMBaseModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import (
|
||||
|
|
@ -901,6 +902,93 @@ class ContentFilterConfigModel(LiteLLMBaseModel):
|
|||
|
||||
MCP_SECURITY_ON_VIOLATION: Final = frozenset({"block", "alert"})
|
||||
|
||||
GuardrailStreamScope = Literal["streaming", "non_streaming", "both"]
|
||||
DEFAULT_GUARDRAIL_STREAM_SCOPE: Final[GuardrailStreamScope] = "both"
|
||||
|
||||
|
||||
class GuardrailEventHooks(str, Enum):
|
||||
pre_call = "pre_call"
|
||||
post_call = "post_call"
|
||||
during_call = "during_call"
|
||||
logging_only = "logging_only"
|
||||
pre_mcp_call = "pre_mcp_call"
|
||||
during_mcp_call = "during_mcp_call"
|
||||
post_mcp_call = "post_mcp_call"
|
||||
realtime_input_transcription = "realtime_input_transcription"
|
||||
|
||||
|
||||
GUARDRAIL_EVENT_HOOK_VALUES: Final = frozenset(member.value for member in GuardrailEventHooks)
|
||||
|
||||
_GUARDRAIL_STREAM_SCOPES: Final[Mapping[str, GuardrailStreamScope]] = MappingProxyType(
|
||||
{
|
||||
"streaming": "streaming",
|
||||
"non_streaming": "non_streaming",
|
||||
"both": "both",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _as_guardrail_stream_scope(value: object) -> GuardrailStreamScope:
|
||||
if not isinstance(value, str):
|
||||
raise ValueError(f"stream_scope values must be strings, got {type(value).__name__}")
|
||||
scope: Final = _GUARDRAIL_STREAM_SCOPES.get(value.lower())
|
||||
if scope is None:
|
||||
raise ValueError(f"stream_scope must be one of both, streaming, non_streaming, got {value!r}")
|
||||
return scope
|
||||
|
||||
|
||||
def _validated_stream_scope_hook(key: object) -> str:
|
||||
if not isinstance(key, str):
|
||||
raise ValueError(f"stream_scope keys must be strings, got {type(key).__name__}")
|
||||
hook: Final = key.lower()
|
||||
if hook not in GUARDRAIL_EVENT_HOOK_VALUES:
|
||||
raise ValueError(
|
||||
f"stream_scope keys must be guardrail modes ({sorted(GUARDRAIL_EVENT_HOOK_VALUES)}), got {key!r}"
|
||||
)
|
||||
return hook
|
||||
|
||||
|
||||
def coerce_stream_scope(value: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return _as_guardrail_stream_scope(value)
|
||||
if isinstance(value, Mapping):
|
||||
scope_map: Final[Mapping[str, object]] = cast(Mapping[str, object], value) # cast-ok: keys validated below
|
||||
return {
|
||||
_validated_stream_scope_hook(key): _as_guardrail_stream_scope(scope) for key, scope in scope_map.items()
|
||||
}
|
||||
raise ValueError(f"stream_scope must be a string or mapping, got {type(value).__name__}")
|
||||
|
||||
|
||||
def stored_stream_scope(value: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None:
|
||||
try:
|
||||
return coerce_stream_scope(value)
|
||||
except ValueError:
|
||||
verbose_logger.warning("Ignoring invalid stored stream_scope value of type %s", type(value).__name__)
|
||||
return None
|
||||
|
||||
|
||||
def with_tolerated_stream_scope(params: Mapping[str, object]) -> dict[str, object]:
|
||||
if "stream_scope" not in params:
|
||||
return dict(params)
|
||||
return {
|
||||
**params,
|
||||
"stream_scope": stored_stream_scope(params["stream_scope"]),
|
||||
}
|
||||
|
||||
|
||||
def runtime_stream_scope(
|
||||
stream_scope: object,
|
||||
) -> tuple[GuardrailStreamScope, MappingProxyType[str, GuardrailStreamScope]]:
|
||||
coerced: Final = coerce_stream_scope(stream_scope)
|
||||
if coerced is None:
|
||||
return DEFAULT_GUARDRAIL_STREAM_SCOPE, MappingProxyType({})
|
||||
if isinstance(coerced, str):
|
||||
return coerced, MappingProxyType({})
|
||||
return DEFAULT_GUARDRAIL_STREAM_SCOPE, MappingProxyType(coerced)
|
||||
|
||||
|
||||
LoggingOnlyScope = Literal["input", "output", "both"]
|
||||
|
||||
|
||||
|
|
@ -1144,6 +1232,21 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
),
|
||||
)
|
||||
|
||||
stream_scope: GuardrailStreamScope | dict[str, GuardrailStreamScope] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Whether this guardrail runs on streaming requests, non-streaming requests, or both. "
|
||||
"A string applies to every configured mode. A map overrides named modes "
|
||||
"(pre_call, during_call, post_call, ...); omitted keys default to both. "
|
||||
"Unset means both, matching historical behavior."
|
||||
),
|
||||
)
|
||||
|
||||
@field_validator("stream_scope", mode="before")
|
||||
@classmethod
|
||||
def normalize_stream_scope(cls, v: object) -> GuardrailStreamScope | dict[str, GuardrailStreamScope] | None:
|
||||
return coerce_stream_scope(v)
|
||||
|
||||
logging_only_scope: LoggingOnlyScope | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
@ -1278,17 +1381,6 @@ class guardrailConfig(TypedDict):
|
|||
guardrails: list[Guardrail]
|
||||
|
||||
|
||||
class GuardrailEventHooks(str, Enum):
|
||||
pre_call = "pre_call"
|
||||
post_call = "post_call"
|
||||
during_call = "during_call"
|
||||
logging_only = "logging_only"
|
||||
pre_mcp_call = "pre_mcp_call"
|
||||
during_mcp_call = "during_mcp_call"
|
||||
post_mcp_call = "post_mcp_call"
|
||||
realtime_input_transcription = "realtime_input_transcription"
|
||||
|
||||
|
||||
class DynamicGuardrailParams(TypedDict):
|
||||
extra_body: ReadOnly[dict[str, object]]
|
||||
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ class PipelineStepResult(LiteLLMBaseModel):
|
|||
"""Result of executing a single pipeline step."""
|
||||
|
||||
guardrail_name: str
|
||||
outcome: Literal["pass", "fail", "error"]
|
||||
outcome: Literal["pass", "fail", "error", "skip"]
|
||||
action_taken: str
|
||||
modified_data: dict[str, Any] | None = None
|
||||
error_detail: str | None = None
|
||||
|
|
|
|||
856
tests/integration/observability/test_guardrail_stream_scope.py
Normal file
856
tests/integration/observability/test_guardrail_stream_scope.py
Normal file
|
|
@ -0,0 +1,856 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, JsonValue, Scenario, eventually
|
||||
from integration._support.database import read_rows, write_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.upstream import (
|
||||
_aws_event_frame, # pyright: ignore[reportPrivateUsage] # project Bedrock event-stream encoder
|
||||
)
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
EndpointScope: TypeAlias = Literal["streaming", "non_streaming"]
|
||||
ModelKind: TypeAlias = Literal["router", "direct"]
|
||||
StreamingAction: TypeAlias = Literal["converse-stream", "invoke-with-response-stream"]
|
||||
NonStreamingAction: TypeAlias = Literal["converse", "invoke"]
|
||||
BedrockAction: TypeAlias = StreamingAction | NonStreamingAction
|
||||
BEDROCK_MODEL_ID: Final = "anthropic.claude-sonnet-5-v1:0"
|
||||
BEDROCK_FALSE_POSITIVE_MODEL_ID: Final = "anthropic.claude-converse-stream-test-v1:0"
|
||||
BEDROCK_EVENT_STREAM: Final = "application/vnd.amazon.eventstream"
|
||||
GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api"
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: Final = "litellm_server_streaming_classification"
|
||||
STREAMING_ACTIONS: Final[tuple[StreamingAction, ...]] = (
|
||||
"converse-stream",
|
||||
"invoke-with-response-stream",
|
||||
)
|
||||
NON_STREAMING_ACTIONS: Final[tuple[NonStreamingAction, ...]] = ("converse", "invoke")
|
||||
SCOPES: Final[tuple[EndpointScope, ...]] = ("streaming", "non_streaming")
|
||||
HOSTILE_CLASSIFICATION_CASES: Final = (
|
||||
pytest.param("is_streaming_request", True, id="boolean-marker"),
|
||||
pytest.param("is_streaming_request", "litellm-server-streaming", id="server-marker-string"),
|
||||
pytest.param("litellm_server_streaming_classification", True, id="classification-field"),
|
||||
)
|
||||
WORKTREE: Final = Path(__file__).resolve().parents[3]
|
||||
LITELLM_PATH: Final = Path(litellm.__file__).resolve()
|
||||
assert LITELLM_PATH.is_relative_to(WORKTREE), (LITELLM_PATH, WORKTREE)
|
||||
print(f"stream_scope repro litellm import: {LITELLM_PATH}") # noqa: T201 # required worktree evidence
|
||||
|
||||
|
||||
def _json(value: object) -> bytes:
|
||||
return json.dumps(value, separators=(",", ":")).encode()
|
||||
|
||||
|
||||
def _strings(value: JsonValue) -> tuple[str, ...]:
|
||||
if isinstance(value, str):
|
||||
return (value,)
|
||||
if isinstance(value, list):
|
||||
return tuple(chain.from_iterable(_strings(item) for item in value))
|
||||
if isinstance(value, dict):
|
||||
return tuple(chain.from_iterable(_strings(item) for item in value.values()))
|
||||
return ()
|
||||
|
||||
|
||||
def _key_names(value: JsonValue) -> tuple[str, ...]:
|
||||
if isinstance(value, dict):
|
||||
return tuple(value) + tuple(chain.from_iterable(_key_names(item) for item in value.values()))
|
||||
if isinstance(value, list):
|
||||
return tuple(chain.from_iterable(_key_names(item) for item in value))
|
||||
return ()
|
||||
|
||||
|
||||
def _marker(body: JsonValue) -> str:
|
||||
return next(value for value in _strings(body) if value.startswith("scope-"))
|
||||
|
||||
|
||||
def _bedrock_converse_stream(marker: str) -> bytes:
|
||||
return b"".join(
|
||||
_aws_event_frame(event_type, payload, marker, marker)
|
||||
for event_type, payload in (
|
||||
("messageStart", {"role": "assistant"}),
|
||||
(
|
||||
"contentBlockDelta",
|
||||
{"delta": {"text": f"scripted Bedrock reply {marker}"}, "contentBlockIndex": 0},
|
||||
),
|
||||
("messageStop", {"stopReason": "end_turn"}),
|
||||
("metadata", {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _invoke_chunk(payload: Mapping[str, JsonValue], marker: str) -> bytes:
|
||||
encoded: Final = base64.b64encode(_json(payload)).decode()
|
||||
return _aws_event_frame("chunk", {"bytes": encoded}, marker, marker)
|
||||
|
||||
|
||||
def _bedrock_invoke_stream(marker: str) -> bytes:
|
||||
events: Final = (
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": f"msg-{marker}",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": BEDROCK_MODEL_ID,
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 11, "output_tokens": 0},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": f"scripted Bedrock reply {marker}"},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4},
|
||||
},
|
||||
{"type": "message_stop"},
|
||||
)
|
||||
return b"".join(_invoke_chunk(event, marker) for event in events)
|
||||
|
||||
|
||||
def _provider(request: Request) -> Reply:
|
||||
if not request.body:
|
||||
return Reply(status=400, body=_json({"error": "empty request body"}))
|
||||
body: Final = JSON_OBJECT.validate_json(request.body)
|
||||
marker: Final = _marker(body)
|
||||
target: Final = request.target.split("?", 1)[0]
|
||||
if target.startswith("/passthrough"):
|
||||
if body.get("stream") is True:
|
||||
streamed_response: Final = _json({"received": body})
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(b"data: " + streamed_response + b"\n\n", b"data: [DONE]\n\n"),
|
||||
)
|
||||
return Reply(body=_json({"received": body}))
|
||||
if target.endswith("/converse-stream"):
|
||||
return Reply(body=_bedrock_converse_stream(marker), content_type=BEDROCK_EVENT_STREAM)
|
||||
if target.endswith("/invoke-with-response-stream"):
|
||||
return Reply(body=_bedrock_invoke_stream(marker), content_type=BEDROCK_EVENT_STREAM)
|
||||
if target.endswith("/converse") or target.endswith("/invoke"):
|
||||
return Reply(body=_json({"output": f"scripted Bedrock reply {marker}"}))
|
||||
if target == "/v1/chat/completions":
|
||||
if body.get("stream") is True:
|
||||
chunk: Final = {
|
||||
"id": f"chatcmpl-{marker}",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant", "content": f"scripted chat reply {marker}"},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
final_chunk: Final = {
|
||||
**chunk,
|
||||
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
|
||||
}
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(
|
||||
b"data: " + _json(chunk) + b"\n\n",
|
||||
b"data: " + _json(final_chunk) + b"\n\n",
|
||||
b"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
return Reply(
|
||||
body=_json(
|
||||
{
|
||||
"id": f"chatcmpl-{marker}",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": f"scripted chat reply {marker}"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
)
|
||||
)
|
||||
return Reply(status=404, body=_json({"error": f"unexpected upstream path: {target}"}))
|
||||
|
||||
|
||||
def _sink(request: Request) -> Reply:
|
||||
assert request.target.endswith(GUARDRAIL_PATH), request.target
|
||||
assert b"scope-" in request.body, request.body.decode()
|
||||
return Reply(body=_json({"action": "NONE"}))
|
||||
|
||||
|
||||
def _rail(name: str, sink: Wire, scope: EndpointScope) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"guardrail_name": name,
|
||||
"litellm_params": {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"stream_scope": scope,
|
||||
"api_base": f"{sink.url}/{name}",
|
||||
"api_key": "synthetic-guardrail-key",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _chat_proxy_config(provider_url: str, guardrails: list[dict[str, JsonValue]]) -> dict[str, object]:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
return {
|
||||
**config,
|
||||
"guardrails": guardrails,
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "scope-invalid-config-chat",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{provider_url}/v1",
|
||||
"api_key": "synthetic-provider-key",
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReproRig:
|
||||
candidate: Gateway
|
||||
direct_candidate: Gateway
|
||||
scenario: Scenario
|
||||
models: Mapping[str, str]
|
||||
rails: Mapping[str, str]
|
||||
provider: Wire
|
||||
sink: Wire
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ReproRig]:
|
||||
with httpx.Client(
|
||||
base_url=os.environ["INTEGRATION_PROXY_URL"],
|
||||
timeout=30,
|
||||
trust_env=False,
|
||||
) as root_client:
|
||||
root_gateway: Final = Gateway(
|
||||
root_client,
|
||||
os.environ.get("INTEGRATION_MASTER_KEY", "sk-integration-master"),
|
||||
os.environ["INTEGRATION_UPSTREAM_URL"],
|
||||
)
|
||||
directory: Final = tmp_path_factory.mktemp("guardrail-stream-scope-repro")
|
||||
with wire_server(_provider) as provider, wire_server(_sink) as sink:
|
||||
rails: Final = MappingProxyType(
|
||||
{
|
||||
"bedrock_streaming": "bedrock_streaming",
|
||||
"bedrock_non_streaming": "bedrock_non_streaming",
|
||||
"chat_streaming": "chat_streaming",
|
||||
"passthrough_streaming": "passthrough_streaming",
|
||||
"passthrough_non_streaming": "passthrough_non_streaming",
|
||||
"passthrough_spoof_streaming": "passthrough_spoof_streaming",
|
||||
}
|
||||
)
|
||||
models: Final = MappingProxyType(
|
||||
{
|
||||
"chat": "scope-chat",
|
||||
"bedrock_router": "scope-bedrock-router",
|
||||
"bedrock_false_positive_router": "scope-bedrock-converse-stream-model",
|
||||
}
|
||||
)
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
_rail(rails["bedrock_streaming"], sink, "streaming"),
|
||||
_rail(rails["bedrock_non_streaming"], sink, "non_streaming"),
|
||||
_rail(rails["chat_streaming"], sink, "streaming"),
|
||||
_rail(rails["passthrough_streaming"], sink, "streaming"),
|
||||
_rail(rails["passthrough_non_streaming"], sink, "non_streaming"),
|
||||
_rail(rails["passthrough_spoof_streaming"], sink, "streaming"),
|
||||
]
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": models["chat"],
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{provider.url}/v1",
|
||||
"api_key": "synthetic-provider-key",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": models["bedrock_router"],
|
||||
"litellm_params": {
|
||||
"model": f"bedrock/{BEDROCK_MODEL_ID}",
|
||||
"api_base": provider.url,
|
||||
"aws_access_key_id": "AKIASYNTHETICSTREAMSCOPE",
|
||||
"aws_secret_access_key": "synthetic-bedrock-secret",
|
||||
"aws_region_name": "us-east-1",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": models["bedrock_false_positive_router"],
|
||||
"litellm_params": {
|
||||
"model": f"bedrock/{BEDROCK_FALSE_POSITIVE_MODEL_ID}",
|
||||
"api_base": provider.url,
|
||||
"aws_access_key_id": "AKIASYNTHETICSTREAMSCOPE",
|
||||
"aws_secret_access_key": "synthetic-bedrock-secret",
|
||||
"aws_region_name": "us-east-1",
|
||||
},
|
||||
},
|
||||
]
|
||||
config["environment_variables"] = {
|
||||
"AWS_BEDROCK_RUNTIME_ENDPOINT": provider.url,
|
||||
"AWS_ACCESS_KEY_ID": "AKIASYNTHETICSTREAMSCOPE",
|
||||
"AWS_SECRET_ACCESS_KEY": "synthetic-bedrock-secret",
|
||||
"AWS_REGION": "us-east-1",
|
||||
"AWS_REGION_NAME": "us-east-1",
|
||||
}
|
||||
config["general_settings"]["pass_through_endpoints"] = [
|
||||
{
|
||||
"path": "/pt-forward",
|
||||
"target": f"{provider.url}/passthrough",
|
||||
"include_subpath": True,
|
||||
},
|
||||
{
|
||||
"path": "/pt-spoof",
|
||||
"target": f"{provider.url}/passthrough",
|
||||
"include_subpath": True,
|
||||
"guardrails": {rails["passthrough_spoof_streaming"]: None},
|
||||
},
|
||||
{
|
||||
"path": "/pt-scope",
|
||||
"target": f"{provider.url}/passthrough",
|
||||
"include_subpath": True,
|
||||
"guardrails": {
|
||||
rails["passthrough_streaming"]: None,
|
||||
rails["passthrough_non_streaming"]: None,
|
||||
},
|
||||
},
|
||||
]
|
||||
config_path: Final = directory / "stream-scope-repro.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy_process(root_gateway, directory, {}, config=config_path, workers=1) as owned:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
direct_config: Final = {
|
||||
**config,
|
||||
"model_list": [
|
||||
*config["model_list"],
|
||||
{
|
||||
"model_name": f"scope-unused-{uuid.uuid4().hex}*",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": f"{provider.url}/v1",
|
||||
"api_key": "synthetic-provider-key",
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
direct_config_path: Final = directory / "stream-scope-direct.yaml"
|
||||
direct_config_path.write_text(yaml.safe_dump(direct_config))
|
||||
with owned_proxy_process(
|
||||
root_gateway,
|
||||
directory,
|
||||
{},
|
||||
config=direct_config_path,
|
||||
workers=1,
|
||||
) as direct:
|
||||
yield ReproRig(
|
||||
owned.gateway,
|
||||
direct.gateway,
|
||||
scenario,
|
||||
models,
|
||||
rails,
|
||||
provider,
|
||||
sink,
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_request(
|
||||
action: BedrockAction,
|
||||
model_path: str,
|
||||
marker: str,
|
||||
) -> tuple[str, dict[str, JsonValue]]:
|
||||
if action in ("converse", "converse-stream"):
|
||||
body: Final = {
|
||||
"messages": [{"role": "user", "content": [{"text": marker}]}],
|
||||
"inferenceConfig": {"maxTokens": 16},
|
||||
}
|
||||
else:
|
||||
body = {
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": marker}]}],
|
||||
}
|
||||
return f"/bedrock/model/{model_path}/{action}", body
|
||||
|
||||
|
||||
def _matching_requests(wire: Wire, marker: str) -> tuple[Request, ...]:
|
||||
return tuple(request for request in wire.drain() if marker.encode() in request.body)
|
||||
|
||||
|
||||
def _rail_scans(rows: Sequence[Request], rail_name: str, marker: str) -> tuple[Request, ...]:
|
||||
return tuple(
|
||||
request for request in rows if request.target.startswith(f"/{rail_name}/") and marker.encode() in request.body
|
||||
)
|
||||
|
||||
|
||||
def _chat_request_with_scans(
|
||||
gateway: Gateway,
|
||||
sink: Wire,
|
||||
model: str,
|
||||
marker: str,
|
||||
streamed: bool,
|
||||
rail_name: str,
|
||||
) -> tuple[httpx.Response, tuple[Request, ...]]:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": marker}],
|
||||
"stream": streamed,
|
||||
"guardrails": [rail_name],
|
||||
},
|
||||
)
|
||||
return response, _rail_scans(sink.drain(), rail_name, marker)
|
||||
|
||||
|
||||
def _spend_row_for_call(call_id: str, content: bytes) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
"SELECT request_id, litellm_call_id, spend, prompt_tokens, completion_tokens, metadata "
|
||||
'FROM "LiteLLM_SpendLogs" WHERE request_id=%s OR litellm_call_id=%s',
|
||||
(call_id, call_id),
|
||||
),
|
||||
lambda values: len(values) >= 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert len(rows) == 1, (call_id, rows, content)
|
||||
assert call_id in (rows[0]["request_id"], rows[0]["litellm_call_id"]), content
|
||||
return rows[0]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_kind", ("router", "direct"))
|
||||
@pytest.mark.parametrize("action", STREAMING_ACTIONS)
|
||||
@pytest.mark.parametrize("scope", SCOPES)
|
||||
def test_bedrock_streaming_actions_run_streaming_scoped_rails(
|
||||
rig: ReproRig,
|
||||
model_kind: ModelKind,
|
||||
action: StreamingAction,
|
||||
scope: EndpointScope,
|
||||
) -> None:
|
||||
marker: Final = f"scope-bedrock-stream-{uuid.uuid4().hex}"
|
||||
model_path: Final = rig.models["bedrock_router"] if model_kind == "router" else BEDROCK_MODEL_ID
|
||||
path, body = _bedrock_request(action, model_path, marker)
|
||||
expected_body: Final = (
|
||||
_bedrock_converse_stream(marker) if action == "converse-stream" else _bedrock_invoke_stream(marker)
|
||||
)
|
||||
call_id: Final = f"stream-scope-{uuid.uuid4().hex}"
|
||||
key: Final = rig.scenario.key(guardrails=[rig.rails[f"bedrock_{scope}"]])
|
||||
candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
path,
|
||||
body,
|
||||
key=key,
|
||||
headers={"x-litellm-call-id": call_id},
|
||||
)
|
||||
assert response.status_code == 200, response.content
|
||||
assert response.headers.get("content-type") == BEDROCK_EVENT_STREAM, dict(response.headers)
|
||||
assert response.content == expected_body, response.content
|
||||
provider_rows: Final = _matching_requests(rig.provider, marker)
|
||||
assert len(provider_rows) == 1, (marker, provider_rows, response.content)
|
||||
provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body)
|
||||
key_names: Final = _key_names(provider_body)
|
||||
assert "is_streaming_request" not in key_names, provider_body
|
||||
assert not tuple(name for name in key_names if name.startswith("litellm_")), provider_body
|
||||
_spend_row_for_call(call_id, response.content)
|
||||
sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails[f"bedrock_{scope}"], marker)
|
||||
expected_scans: Final = int(scope == "streaming")
|
||||
assert len(sink_rows) == expected_scans, (marker, model_kind, action, scope, sink_rows, response.content)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_kind", ("router", "direct"))
|
||||
@pytest.mark.parametrize("action", NON_STREAMING_ACTIONS)
|
||||
@pytest.mark.parametrize("scope", SCOPES)
|
||||
def test_bedrock_non_streaming_actions_run_non_streaming_scoped_rails(
|
||||
rig: ReproRig,
|
||||
model_kind: ModelKind,
|
||||
action: NonStreamingAction,
|
||||
scope: EndpointScope,
|
||||
) -> None:
|
||||
marker: Final = f"scope-bedrock-nonstream-{uuid.uuid4().hex}"
|
||||
model_path: Final = rig.models["bedrock_router"] if model_kind == "router" else BEDROCK_MODEL_ID
|
||||
path, body = _bedrock_request(action, model_path, marker)
|
||||
key: Final = rig.scenario.key(guardrails=[rig.rails[f"bedrock_{scope}"]])
|
||||
candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate
|
||||
response: Final = candidate.request("POST", path, body, key=key)
|
||||
assert response.status_code == 200, response.text
|
||||
assert marker in response.text, response.text
|
||||
provider_rows: Final = _matching_requests(rig.provider, marker)
|
||||
assert len(provider_rows) == 1, (marker, provider_rows, response.text)
|
||||
provider_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body)
|
||||
key_names: Final = _key_names(provider_body)
|
||||
assert "is_streaming_request" not in key_names, provider_body
|
||||
assert not tuple(name for name in key_names if name.startswith("litellm_")), provider_body
|
||||
sink_rows: Final = _rail_scans(rig.sink.drain(), rig.rails[f"bedrock_{scope}"], marker)
|
||||
expected_scans: Final = int(scope == "non_streaming")
|
||||
assert len(sink_rows) == expected_scans, (marker, model_kind, action, scope, sink_rows, response.text)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_kind", ("router", "direct"))
|
||||
def test_bedrock_model_id_streaming_action_text_on_converse_is_non_streaming(
|
||||
rig: ReproRig,
|
||||
model_kind: ModelKind,
|
||||
) -> None:
|
||||
marker: Final = f"scope-bedrock-converse-model-{uuid.uuid4().hex}"
|
||||
model_path: Final = (
|
||||
rig.models["bedrock_false_positive_router"] if model_kind == "router" else BEDROCK_FALSE_POSITIVE_MODEL_ID
|
||||
)
|
||||
path, body = _bedrock_request("converse", model_path, marker)
|
||||
call_id: Final = f"stream-scope-bedrock-{uuid.uuid4().hex}"
|
||||
key: Final = rig.scenario.key(
|
||||
guardrails=[
|
||||
rig.rails["bedrock_streaming"],
|
||||
rig.rails["bedrock_non_streaming"],
|
||||
]
|
||||
)
|
||||
candidate: Final = rig.candidate if model_kind == "router" else rig.direct_candidate
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
path,
|
||||
body,
|
||||
key=key,
|
||||
headers={"x-litellm-call-id": call_id},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert marker in response.text, response.text
|
||||
provider_rows: Final = _matching_requests(rig.provider, marker)
|
||||
assert len(provider_rows) == 1, (marker, model_kind, provider_rows, response.text)
|
||||
sink_rows: Final = rig.sink.drain()
|
||||
streaming_rows: Final = _rail_scans(sink_rows, rig.rails["bedrock_streaming"], marker)
|
||||
non_streaming_rows: Final = _rail_scans(sink_rows, rig.rails["bedrock_non_streaming"], marker)
|
||||
assert streaming_rows == (), (marker, model_kind, streaming_rows, response.text)
|
||||
assert len(non_streaming_rows) == 1, (marker, model_kind, non_streaming_rows, response.text)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("streamed", (False, True), ids=("stream-absent", "stream-true"))
|
||||
def test_configured_passthrough_forwards_caller_is_streaming_request_field(
|
||||
rig: ReproRig,
|
||||
streamed: bool,
|
||||
) -> None:
|
||||
marker: Final = f"scope-passthrough-forward-{uuid.uuid4().hex}"
|
||||
caller_value: Final = f"caller-{uuid.uuid4().hex}"
|
||||
body: Final = {
|
||||
"marker": marker,
|
||||
"is_streaming_request": caller_value,
|
||||
**({"stream": True} if streamed else {}),
|
||||
}
|
||||
response: Final = rig.candidate.request("POST", "/pt-forward", body)
|
||||
assert response.status_code == 200, response.text
|
||||
provider_rows: Final = _matching_requests(rig.provider, marker)
|
||||
assert len(provider_rows) == 1, (marker, provider_rows, response.text)
|
||||
upstream_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body)
|
||||
assert upstream_body.get("is_streaming_request") == caller_value, (
|
||||
marker,
|
||||
caller_value,
|
||||
upstream_body,
|
||||
response.text,
|
||||
)
|
||||
assert upstream_body == body, (marker, body, upstream_body, response.text)
|
||||
if streamed:
|
||||
assert response.headers.get("content-type", "").lower().startswith("text/event-stream"), dict(response.headers)
|
||||
event_body: Final = response.text.removeprefix("data: ").split("\n", maxsplit=1)[0]
|
||||
response_body: Final = JSON_OBJECT.validate_json(event_body)
|
||||
else:
|
||||
response_body = JSON_OBJECT.validate_json(response.content)
|
||||
assert response_body == {"received": body}, response.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("hostile_field", "hostile_value"), HOSTILE_CLASSIFICATION_CASES)
|
||||
def test_client_cannot_spoof_server_stream_classification(
|
||||
rig: ReproRig,
|
||||
hostile_field: str,
|
||||
hostile_value: JsonValue,
|
||||
) -> None:
|
||||
chat_marker: Final = f"scope-chat-spoof-{uuid.uuid4().hex}"
|
||||
chat_body: Final = {
|
||||
"model": rig.models["chat"],
|
||||
"messages": [{"role": "user", "content": chat_marker}],
|
||||
"stream": False,
|
||||
hostile_field: hostile_value,
|
||||
}
|
||||
chat_key: Final = rig.scenario.key(guardrails=[rig.rails["chat_streaming"]])
|
||||
chat_response: Final = rig.candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
chat_body,
|
||||
key=chat_key,
|
||||
)
|
||||
chat_provider_rows: Final = _matching_requests(rig.provider, chat_marker)
|
||||
chat_upstream_body: Final = chat_provider_rows[0].body.decode() if chat_provider_rows else "<no upstream request>"
|
||||
print( # noqa: T201 # required chat hostile-body observation
|
||||
f"chat hostile {hostile_field}={hostile_value!r}: "
|
||||
f"status={chat_response.status_code}, response={chat_response.text!r}, upstream={chat_upstream_body}"
|
||||
)
|
||||
assert chat_response.status_code == 200, chat_response.text
|
||||
assert len(chat_provider_rows) == 1, (chat_marker, chat_provider_rows, chat_response.text)
|
||||
assert JSON_OBJECT.validate_json(chat_provider_rows[0].body) == {
|
||||
"messages": [{"role": "user", "content": chat_marker}],
|
||||
"model": "gpt-4o-mini",
|
||||
hostile_field: hostile_value,
|
||||
}, (hostile_field, hostile_value, chat_upstream_body)
|
||||
assert _rail_scans(rig.sink.drain(), rig.rails["chat_streaming"], chat_marker) == (), (
|
||||
chat_marker,
|
||||
hostile_field,
|
||||
hostile_value,
|
||||
chat_response.text,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("hostile_field", "hostile_value"), HOSTILE_CLASSIFICATION_CASES)
|
||||
def test_configured_passthrough_cannot_spoof_server_stream_classification(
|
||||
rig: ReproRig,
|
||||
hostile_field: str,
|
||||
hostile_value: JsonValue,
|
||||
) -> None:
|
||||
passthrough_marker: Final = f"scope-passthrough-spoof-{uuid.uuid4().hex}"
|
||||
passthrough_body: Final = {
|
||||
"marker": passthrough_marker,
|
||||
"stream": False,
|
||||
hostile_field: hostile_value,
|
||||
}
|
||||
passthrough_response: Final = rig.candidate.request("POST", "/pt-spoof", passthrough_body)
|
||||
assert passthrough_response.status_code == 200, passthrough_response.text
|
||||
passthrough_provider_rows: Final = _matching_requests(rig.provider, passthrough_marker)
|
||||
assert len(passthrough_provider_rows) == 1, (
|
||||
passthrough_marker,
|
||||
passthrough_provider_rows,
|
||||
passthrough_response.text,
|
||||
)
|
||||
assert _rail_scans(rig.sink.drain(), rig.rails["passthrough_spoof_streaming"], passthrough_marker) == (), (
|
||||
passthrough_marker,
|
||||
hostile_field,
|
||||
hostile_value,
|
||||
passthrough_response.text,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("streamed", "request_fields"),
|
||||
((True, {"stream": True}), (False, {"stream": False}), (False, {})),
|
||||
ids=("stream-true", "stream-false", "stream-absent"),
|
||||
)
|
||||
def test_passthrough_scope_follows_proxy_stream_decision(
|
||||
rig: ReproRig,
|
||||
streamed: bool,
|
||||
request_fields: dict[str, JsonValue],
|
||||
) -> None:
|
||||
marker: Final = f"scope-passthrough-scope-{uuid.uuid4().hex}"
|
||||
body: Final = {"marker": marker, **request_fields}
|
||||
response: Final = rig.candidate.request("POST", "/pt-scope", body)
|
||||
assert response.status_code == 200, response.text
|
||||
if streamed:
|
||||
expected_frame: Final = f"data: {_json({'received': body}).decode()}\n\ndata: [DONE]\n\n"
|
||||
assert response.headers.get("content-type", "").lower().startswith("text/event-stream"), dict(response.headers)
|
||||
assert response.text == expected_frame, response.text
|
||||
assert response.headers.get("transfer-encoding", "").lower() == "chunked", dict(response.headers)
|
||||
else:
|
||||
response_body: Final = JSON_OBJECT.validate_json(response.content)
|
||||
assert response_body == {"received": body}, response.text
|
||||
assert "content-length" in response.headers, dict(response.headers)
|
||||
provider_rows: Final = _matching_requests(rig.provider, marker)
|
||||
assert len(provider_rows) == 1, (marker, provider_rows, response.text)
|
||||
upstream_body: Final = JSON_OBJECT.validate_json(provider_rows[0].body)
|
||||
assert upstream_body == body, (marker, body, upstream_body, response.text)
|
||||
assert SERVER_STREAMING_CLASSIFICATION_KEY not in upstream_body, upstream_body
|
||||
sink_rows: Final = rig.sink.drain()
|
||||
streaming_rows: Final = _rail_scans(sink_rows, rig.rails["passthrough_streaming"], marker)
|
||||
non_streaming_rows: Final = _rail_scans(sink_rows, rig.rails["passthrough_non_streaming"], marker)
|
||||
assert len(streaming_rows) == int(streamed), (marker, streamed, streaming_rows, response.text)
|
||||
assert len(non_streaming_rows) == int(not streamed), (marker, streamed, non_streaming_rows, response.text)
|
||||
|
||||
|
||||
def test_invalid_yaml_stream_scope_keeps_rail_running_on_both_shapes(rig: ReproRig, tmp_path: Path) -> None:
|
||||
name: Final = f"scope-invalid-yaml-{uuid.uuid4().hex}"
|
||||
invalid_rail: Final = {
|
||||
"guardrail_name": name,
|
||||
"litellm_params": {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"stream_scope": "sometimes",
|
||||
"api_base": f"{rig.sink.url}/{name}",
|
||||
"api_key": "synthetic-guardrail-key",
|
||||
},
|
||||
}
|
||||
config: Final = _chat_proxy_config(rig.provider.url, [invalid_rail])
|
||||
config_path: Final = tmp_path / "invalid-stream-scope.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned:
|
||||
markers: Final = tuple(f"scope-invalid-yaml-{int(streamed)}-{uuid.uuid4().hex}" for streamed in (False, True))
|
||||
observations: Final = tuple(
|
||||
_chat_request_with_scans(
|
||||
owned.gateway,
|
||||
rig.sink,
|
||||
"scope-invalid-config-chat",
|
||||
marker,
|
||||
streamed,
|
||||
name,
|
||||
)
|
||||
for streamed, marker in zip((False, True), markers)
|
||||
)
|
||||
assert tuple(response.status_code for response, _ in observations) == (200, 200), tuple(
|
||||
response.text for response, _ in observations
|
||||
)
|
||||
assert tuple(len(sink_rows) for _, sink_rows in observations) == (1, 1), (markers, observations)
|
||||
listed: Final = owned.gateway.request("GET", "/guardrails/list")
|
||||
assert listed.status_code == 200, listed.text
|
||||
list_payload: Final = JSON_OBJECT.validate_json(listed.content)
|
||||
listed_guardrails: Final = list_payload.get("guardrails")
|
||||
assert isinstance(listed_guardrails, list), list_payload
|
||||
listed_rail: Final = next(
|
||||
(row for row in listed_guardrails if isinstance(row, dict) and row.get("guardrail_name") == name),
|
||||
None,
|
||||
)
|
||||
assert isinstance(listed_rail, dict), list_payload
|
||||
listed_params: Final = listed_rail.get("litellm_params")
|
||||
assert isinstance(listed_params, dict), listed_rail
|
||||
assert listed_params.get("stream_scope") is None, listed_rail
|
||||
|
||||
|
||||
def test_persisted_invalid_stream_scope_row_stays_readable_and_enforced(
|
||||
rig: ReproRig,
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
guardrail_id: Final = str(uuid.uuid4())
|
||||
guardrail_name: Final = f"scope-invalid-persisted-{uuid.uuid4().hex}"
|
||||
params: Final = {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": f"{rig.sink.url}/{guardrail_name}",
|
||||
"api_key": "synthetic-guardrail-key",
|
||||
"stream_scope": "sometimes",
|
||||
}
|
||||
database_url: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL") or os.environ["DATABASE_URL"]
|
||||
write_rows(
|
||||
'INSERT INTO "LiteLLM_GuardrailsTable" '
|
||||
'("guardrail_id", "guardrail_name", "litellm_params", "created_at", "updated_at") '
|
||||
"VALUES (%s, %s, %s::jsonb, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)",
|
||||
(guardrail_id, guardrail_name, json.dumps(params)),
|
||||
database_url=database_url,
|
||||
)
|
||||
try:
|
||||
config: Final = _chat_proxy_config(rig.provider.url, [])
|
||||
config_path: Final = tmp_path / "persisted-invalid-stream-scope.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy_process(rig.candidate, tmp_path, {}, config=config_path, workers=1) as owned:
|
||||
info: Final = eventually(
|
||||
lambda: owned.gateway.request("GET", f"/guardrails/{guardrail_id}/info"),
|
||||
lambda response: response.status_code != 404,
|
||||
seconds=20,
|
||||
)
|
||||
listed: Final = owned.gateway.request("GET", "/v2/guardrails/list")
|
||||
list_payload: Final = JSON_OBJECT.validate_json(listed.content) if listed.status_code == 200 else {}
|
||||
listed_guardrails: Final = list_payload.get("guardrails")
|
||||
includes_row: Final = isinstance(listed_guardrails, list) and any(
|
||||
isinstance(row, dict) and row.get("guardrail_id") == guardrail_id for row in listed_guardrails
|
||||
)
|
||||
markers: Final = (
|
||||
f"scope-invalid-persisted-0-{uuid.uuid4().hex}",
|
||||
f"scope-invalid-persisted-1-{uuid.uuid4().hex}",
|
||||
)
|
||||
observations: Final = tuple(
|
||||
_chat_request_with_scans(
|
||||
owned.gateway,
|
||||
rig.sink,
|
||||
"scope-invalid-config-chat",
|
||||
marker,
|
||||
streamed,
|
||||
guardrail_name,
|
||||
)
|
||||
for streamed, marker in zip((False, True), markers)
|
||||
)
|
||||
assert (
|
||||
info.status_code == 200
|
||||
and listed.status_code == 200
|
||||
and includes_row
|
||||
and tuple(response.status_code for response, _ in observations) == (200, 200)
|
||||
and tuple(len(sink_rows) for _, sink_rows in observations) == (1, 1)
|
||||
), {
|
||||
"info": (info.status_code, info.text),
|
||||
"list": (listed.status_code, listed.text),
|
||||
"includes_row": includes_row,
|
||||
"responses": tuple((response.status_code, response.text) for response, _ in observations),
|
||||
"scan_counts": tuple(len(sink_rows) for _, sink_rows in observations),
|
||||
}
|
||||
finally:
|
||||
write_rows(
|
||||
'DELETE FROM "LiteLLM_GuardrailsTable" WHERE guardrail_id=%s',
|
||||
(guardrail_id,),
|
||||
database_url=database_url,
|
||||
)
|
||||
|
||||
|
||||
def test_management_rejects_invalid_stream_scope(rig: ReproRig) -> None:
|
||||
name: Final = f"scope-invalid-management-{uuid.uuid4().hex}"
|
||||
invalid_params: Final = {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"default_on": False,
|
||||
"api_base": f"{rig.sink.url}/{name}",
|
||||
"api_key": "synthetic-guardrail-key",
|
||||
"stream_scope": "sometimes",
|
||||
}
|
||||
created_invalid: Final = rig.candidate.request(
|
||||
"POST",
|
||||
"/guardrails",
|
||||
{"guardrail": {"guardrail_name": name, "litellm_params": invalid_params}},
|
||||
)
|
||||
assert created_invalid.status_code == 422, created_invalid.text
|
||||
|
||||
valid_params: Final = {**invalid_params, "stream_scope": "both"}
|
||||
created: Final = rig.candidate.request(
|
||||
"POST",
|
||||
"/guardrails",
|
||||
{"guardrail": {"guardrail_name": name, "litellm_params": valid_params}},
|
||||
)
|
||||
assert created.status_code == 200, created.text
|
||||
guardrail_id: Final = str(created.json()["guardrail_id"])
|
||||
try:
|
||||
put_response: Final = rig.candidate.request(
|
||||
"PUT",
|
||||
f"/guardrails/{guardrail_id}",
|
||||
{"guardrail": {"guardrail_name": name, "litellm_params": invalid_params}},
|
||||
)
|
||||
patch_response: Final = rig.candidate.request(
|
||||
"PATCH",
|
||||
f"/guardrails/{guardrail_id}",
|
||||
{"litellm_params": {"stream_scope": "sometimes"}},
|
||||
)
|
||||
assert put_response.status_code == 422, put_response.text
|
||||
assert patch_response.status_code == 422, patch_response.text
|
||||
finally:
|
||||
deleted: Final = rig.candidate.request("DELETE", f"/guardrails/{guardrail_id}")
|
||||
assert deleted.status_code == 200, deleted.text
|
||||
1079
tests/integration/observability/test_guardrail_stream_scope_chaos.py
Normal file
1079
tests/integration/observability/test_guardrail_stream_scope_chaos.py
Normal file
File diff suppressed because it is too large
Load diff
File diff suppressed because it is too large
Load diff
|
|
@ -2936,8 +2936,11 @@ async def test_cancellation_delivers_termination_over_tcp(
|
|||
|
||||
listener: Final = await asyncio.start_server(handle_connection, "127.0.0.1", 0)
|
||||
port: Final = listener.sockets[0].getsockname()[1]
|
||||
client_timeout: Final = 2 if cancel_mode == "read_timeout" else 30
|
||||
client: Final = MCPClient(
|
||||
server_url=f"http://127.0.0.1:{port}/mcp", protocol_version=protocol_version, timeout=2 if cancel_mode == "read_timeout" else 30
|
||||
server_url=f"http://127.0.0.1:{port}/mcp",
|
||||
protocol_version=protocol_version,
|
||||
timeout=client_timeout,
|
||||
)
|
||||
|
||||
async def calls():
|
||||
|
|
|
|||
|
|
@ -1,5 +1,9 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import datetime as dt
|
||||
import json
|
||||
import pickle
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, Literal, Optional
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
|
@ -8,7 +12,10 @@ import pytest
|
|||
from litellm.integrations.custom_guardrail import (
|
||||
DEFAULT_ADVISORY_MESSAGE,
|
||||
CustomGuardrail,
|
||||
_request_is_streaming,
|
||||
guardrail_request_data_with_streaming,
|
||||
log_guardrail_information,
|
||||
without_server_streaming_classification,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
|
||||
|
|
@ -531,6 +538,233 @@ class TestCustomGuardrailShouldRunGuardrail:
|
|||
assert always_on.should_run_guardrail(data=forged, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
|
||||
_STREAM_SCOPE_HOOKS: Final = (
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.during_call,
|
||||
GuardrailEventHooks.post_call,
|
||||
)
|
||||
|
||||
|
||||
class TestCustomGuardrailStreamScope:
|
||||
@pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS)
|
||||
@pytest.mark.parametrize("stream", [True, False])
|
||||
@pytest.mark.parametrize("stream_scope", [None, "both"])
|
||||
def test_both_and_omitted_run_on_streaming_and_non_streaming(
|
||||
self,
|
||||
event_type: GuardrailEventHooks,
|
||||
stream: bool,
|
||||
stream_scope: str | None,
|
||||
):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=event_type,
|
||||
stream_scope=stream_scope,
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": stream}, event_type) is True
|
||||
|
||||
@pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS)
|
||||
def test_scalar_streaming_skips_non_streaming(self, event_type: GuardrailEventHooks):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=event_type,
|
||||
stream_scope="streaming",
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": True}, event_type) is True
|
||||
assert guardrail.should_run_guardrail({"stream": False}, event_type) is False
|
||||
assert guardrail.should_run_guardrail({}, event_type) is False
|
||||
|
||||
@pytest.mark.parametrize("event_type", _STREAM_SCOPE_HOOKS)
|
||||
def test_scalar_non_streaming_skips_streaming(self, event_type: GuardrailEventHooks):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=event_type,
|
||||
stream_scope="non_streaming",
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": False}, event_type) is True
|
||||
assert guardrail.should_run_guardrail({}, event_type) is True
|
||||
assert guardrail.should_run_guardrail({"stream": True}, event_type) is False
|
||||
|
||||
def test_per_mode_map_applies_to_named_hooks_only(self):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call],
|
||||
stream_scope={"pre_call": "both", "post_call": "streaming"},
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True
|
||||
assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.pre_call) is True
|
||||
assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True
|
||||
assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False
|
||||
|
||||
def test_default_on_early_return_still_honors_stream_scope(self):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
stream_scope="non_streaming",
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is True
|
||||
assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is False
|
||||
|
||||
def test_apply_stream_scope_overwrites_constructor_default(self):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="scoped",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is True
|
||||
guardrail.apply_stream_scope("streaming")
|
||||
assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True
|
||||
assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False
|
||||
|
||||
def test_direct_constructor_normalizes_mixed_case_map_keys(self):
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope={"Pre_Call": "streaming"},
|
||||
)
|
||||
assert guardrail.should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True
|
||||
assert guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.pre_call) is False
|
||||
|
||||
def test_realtime_transcription_counts_as_streaming(self):
|
||||
streaming_only = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.realtime_input_transcription,
|
||||
stream_scope="streaming",
|
||||
)
|
||||
assert (
|
||||
streaming_only.should_run_guardrail(
|
||||
{"litellm_metadata": {}}, GuardrailEventHooks.realtime_input_transcription
|
||||
)
|
||||
is True
|
||||
)
|
||||
non_streaming_only = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.realtime_input_transcription,
|
||||
stream_scope="non_streaming",
|
||||
)
|
||||
assert (
|
||||
non_streaming_only.should_run_guardrail(
|
||||
{"litellm_metadata": {}}, GuardrailEventHooks.realtime_input_transcription
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_path_defined_streaming_classification_cannot_be_spoofed(self):
|
||||
generate_content_body: Final = {"contents": [{"parts": [{"text": "hi"}]}]}
|
||||
streaming_only = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="streaming",
|
||||
)
|
||||
assert streaming_only.should_run_guardrail(generate_content_body, GuardrailEventHooks.pre_call) is False
|
||||
assert (
|
||||
streaming_only.should_run_guardrail(
|
||||
{**generate_content_body, "is_streaming_request": True},
|
||||
GuardrailEventHooks.pre_call,
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
streaming_only.should_run_guardrail(
|
||||
{**generate_content_body, "is_streaming_request": "litellm-server-streaming"},
|
||||
GuardrailEventHooks.pre_call,
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
streaming_only.should_run_guardrail(
|
||||
{**generate_content_body, "litellm_server_streaming_classification": True},
|
||||
GuardrailEventHooks.pre_call,
|
||||
)
|
||||
is False
|
||||
)
|
||||
server_streaming_data: Final = guardrail_request_data_with_streaming(
|
||||
generate_content_body,
|
||||
is_streaming=True,
|
||||
)
|
||||
assert (
|
||||
streaming_only.should_run_guardrail(
|
||||
server_streaming_data,
|
||||
GuardrailEventHooks.pre_call,
|
||||
)
|
||||
is True
|
||||
)
|
||||
non_streaming_only = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="non_streaming",
|
||||
)
|
||||
assert (
|
||||
non_streaming_only.should_run_guardrail(
|
||||
server_streaming_data,
|
||||
GuardrailEventHooks.pre_call,
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
def test_streaming_classification_is_json_serializable_without_spoofing(self):
|
||||
d: Final = guardrail_request_data_with_streaming({}, is_streaming=True)
|
||||
serialized: Final = json.dumps(d)
|
||||
assert _request_is_streaming(d) is True
|
||||
|
||||
round_tripped: Final = json.loads(serialized)
|
||||
assert _request_is_streaming(round_tripped) is False
|
||||
assert round_tripped["litellm_server_streaming_classification"] == "litellm-server-streaming"
|
||||
assert isinstance(round_tripped["litellm_server_streaming_classification"], str)
|
||||
assert "litellm_server_streaming_classification" not in without_server_streaming_classification(round_tripped)
|
||||
|
||||
def test_streaming_classification_preserves_caller_fields_and_removes_only_server_marker(self):
|
||||
caller_data: Final = {
|
||||
"contents": [{"parts": [{"text": "hi"}]}],
|
||||
"is_streaming_request": "caller-value",
|
||||
"litellm_server_streaming_classification": True,
|
||||
}
|
||||
non_streaming_data: Final = guardrail_request_data_with_streaming(caller_data, is_streaming=False)
|
||||
server_streaming_data: Final = guardrail_request_data_with_streaming(caller_data, is_streaming=True)
|
||||
|
||||
assert non_streaming_data is not caller_data
|
||||
assert server_streaming_data is not caller_data
|
||||
assert non_streaming_data == caller_data
|
||||
assert server_streaming_data["is_streaming_request"] == "caller-value"
|
||||
assert server_streaming_data["litellm_server_streaming_classification"] is not True
|
||||
assert without_server_streaming_classification(caller_data) == caller_data
|
||||
assert "litellm_server_streaming_classification" not in without_server_streaming_classification(
|
||||
server_streaming_data
|
||||
)
|
||||
|
||||
def test_server_streaming_classification_survives_scan_raw_request_snapshot(self):
|
||||
from litellm.litellm_core_utils.core_helpers import independent_snapshot
|
||||
|
||||
generate_content_body: Final = {"contents": [{"parts": [{"text": "hi"}]}]}
|
||||
snapshot: Final = independent_snapshot(
|
||||
guardrail_request_data_with_streaming(generate_content_body, is_streaming=True)
|
||||
)
|
||||
streaming_only = CustomGuardrail(
|
||||
guardrail_name="test_guardrail",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="streaming",
|
||||
scan_raw_request=True,
|
||||
)
|
||||
assert streaming_only.should_run_guardrail(snapshot, GuardrailEventHooks.pre_call) is True
|
||||
assert (
|
||||
streaming_only.should_run_guardrail(
|
||||
independent_snapshot({**generate_content_body, "is_streaming_request": True}),
|
||||
GuardrailEventHooks.pre_call,
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
class TestApplyGuardrailCheck:
|
||||
def test_apply_guardrail_check_only_on_direct_implementation(self):
|
||||
"""
|
||||
|
|
@ -3307,3 +3541,113 @@ class TestCustomGuardrailTimeout:
|
|||
)
|
||||
|
||||
assert guardrail.timeout == 7.0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream_scope", [None, "both", {"post_call": "streaming"}])
|
||||
def test_guardrail_survives_deepcopy_and_pickle_with_its_stream_scope(stream_scope):
|
||||
guardrail: Final = CustomGuardrail(
|
||||
guardrail_name="copyable",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
stream_scope=stream_scope,
|
||||
)
|
||||
expected: Final = guardrail.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call)
|
||||
for clone in (copy.deepcopy(guardrail), pickle.loads(pickle.dumps(guardrail))):
|
||||
assert dict(clone.stream_scope_by_hook) == dict(guardrail.stream_scope_by_hook)
|
||||
assert clone.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is expected
|
||||
|
||||
|
||||
class _LockHoldingGuardrail(CustomGuardrail):
|
||||
def __init__(self, **kwargs):
|
||||
self.lock = threading.Lock()
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def __getstate__(self):
|
||||
state: Final = dict(self.__dict__)
|
||||
state.pop("lock")
|
||||
return state
|
||||
|
||||
def __setstate__(self, state):
|
||||
self.__dict__.update(state)
|
||||
self.lock = threading.Lock()
|
||||
|
||||
|
||||
class _SetstateOnlyGuardrail(CustomGuardrail):
|
||||
def __setstate__(self, state):
|
||||
self.__dict__.update(state)
|
||||
self.restored = True
|
||||
|
||||
|
||||
class _SlotsGuardrail(CustomGuardrail):
|
||||
__slots__ = ("vendor_client",)
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.vendor_client = "vendor-client-object"
|
||||
|
||||
|
||||
def _subclass_guardrails():
|
||||
kwargs: Final = {
|
||||
"guardrail_name": "vendor",
|
||||
"default_on": True,
|
||||
"event_hook": GuardrailEventHooks.post_call,
|
||||
"stream_scope": {"post_call": "streaming"},
|
||||
}
|
||||
return (
|
||||
pytest.param(_LockHoldingGuardrail(**kwargs), id="dict-getstate"),
|
||||
pytest.param(_SetstateOnlyGuardrail(**kwargs), id="setstate-only"),
|
||||
pytest.param(_SlotsGuardrail(**kwargs), id="slots"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("guardrail", _subclass_guardrails())
|
||||
@pytest.mark.parametrize(
|
||||
"cloner",
|
||||
[copy.copy, copy.deepcopy, lambda g: pickle.loads(pickle.dumps(g))],
|
||||
ids=["copy", "deepcopy", "pickle"],
|
||||
)
|
||||
def test_out_of_tree_guardrail_subclasses_survive_copy_and_pickle(guardrail, cloner):
|
||||
clone = cloner(guardrail)
|
||||
|
||||
if isinstance(clone, _LockHoldingGuardrail):
|
||||
assert isinstance(clone.lock, type(threading.Lock()))
|
||||
elif isinstance(clone, _SetstateOnlyGuardrail):
|
||||
assert clone.restored is True
|
||||
else:
|
||||
assert clone.vendor_client == "vendor-client-object"
|
||||
assert clone.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False
|
||||
assert clone.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True
|
||||
|
||||
|
||||
def test_router_constructs_with_a_dict_getstate_guardrail_in_deployment_callbacks():
|
||||
import litellm
|
||||
|
||||
litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "m",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-5.4-mini",
|
||||
"api_key": "sk-test",
|
||||
"callbacks": [
|
||||
_LockHoldingGuardrail(
|
||||
guardrail_name="vendor",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
stream_scope={"post_call": "streaming"},
|
||||
)
|
||||
],
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_subclass_that_skips_super_init_still_runs_with_default_scope():
|
||||
class NoSuperInit(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
self.guardrail_name = "no-super"
|
||||
self.event_hook = None
|
||||
self.default_on = True
|
||||
|
||||
assert NoSuperInit().should_run_guardrail({"stream": True}, GuardrailEventHooks.pre_call) is True
|
||||
|
|
|
|||
|
|
@ -10,9 +10,14 @@ from datetime import datetime
|
|||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector
|
||||
from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig
|
||||
from litellm.llms.bedrock.passthrough.transformation import (
|
||||
BedrockPassthroughConfig,
|
||||
is_bedrock_streaming_endpoint,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
CONVERSE_MODEL = "anthropic.claude-sonnet-4-5-20250929-v1:0"
|
||||
|
|
@ -20,6 +25,22 @@ CONVERSE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/converse-stream"
|
|||
INVOKE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/invoke-with-response-stream"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("endpoint", "expected"),
|
||||
[
|
||||
("converse-stream", True),
|
||||
("invoke-with-response-stream", True),
|
||||
("converse", False),
|
||||
("invoke", False),
|
||||
("model/my-converse-stream-model/converse", False),
|
||||
("model/x/converse-stream?foo=1", True),
|
||||
("model/x/converse-stream/", True),
|
||||
],
|
||||
)
|
||||
def test_is_bedrock_streaming_endpoint_matches_final_action_segment(endpoint: str, expected: bool) -> None:
|
||||
assert is_bedrock_streaming_endpoint(endpoint) is expected
|
||||
|
||||
|
||||
def test_bedrock_passthrough_get_complete_url_default_endpoint():
|
||||
"""Test get_complete_url with default AWS endpoint (no override)"""
|
||||
config = BedrockPassthroughConfig()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,52 @@
|
|||
import json
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER
|
||||
from litellm.llms.pass_through.guardrail_translation.handler import PassThroughEndpointHandler
|
||||
|
||||
|
||||
def test_full_payload_guardrail_text_excludes_the_server_streaming_marker():
|
||||
body: Final = {"model": "m", "stream": True, "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail(
|
||||
{**body, SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER},
|
||||
None,
|
||||
)
|
||||
|
||||
assert json.loads(text) == body, text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"marker",
|
||||
[
|
||||
SERVER_STREAMING_CLASSIFICATION_MARKER,
|
||||
json.loads(json.dumps(SERVER_STREAMING_CLASSIFICATION_MARKER)),
|
||||
],
|
||||
ids=["enum", "json-string"],
|
||||
)
|
||||
def test_full_payload_guardrail_text_scans_caller_value_but_not_the_marker(marker: str):
|
||||
text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail(
|
||||
{
|
||||
"model": "m",
|
||||
"stream": True,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: "BLOCKME caller content",
|
||||
},
|
||||
None,
|
||||
)
|
||||
|
||||
assert "BLOCKME caller content" in text, text
|
||||
|
||||
marker_text: Final = PassThroughEndpointHandler()._extract_text_for_guardrail(
|
||||
{
|
||||
"model": "m",
|
||||
"stream": True,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: marker,
|
||||
},
|
||||
None,
|
||||
)
|
||||
|
||||
assert SERVER_STREAMING_CLASSIFICATION_KEY not in json.loads(marker_text), marker_text
|
||||
|
|
@ -464,7 +464,7 @@ async def test_sse_mcp_handler_mock():
|
|||
|
||||
mock_sse = MagicMock()
|
||||
mock_sse.connect_sse.side_effect = connect_sse
|
||||
run = AsyncMock()
|
||||
serve = AsyncMock()
|
||||
|
||||
# Mock scope, receive, send with proper ASGI scope format
|
||||
mock_scope = {
|
||||
|
|
@ -489,7 +489,7 @@ async def test_sse_mcp_handler_mock():
|
|||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy._experimental.mcp_server.server.serve_loop", run),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.serve_loop", serve),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
|
|
@ -505,13 +505,21 @@ async def test_sse_mcp_handler_mock():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers",
|
||||
new=AsyncMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth",
|
||||
new=AsyncMock(),
|
||||
),
|
||||
):
|
||||
from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp
|
||||
|
||||
# Call the handler
|
||||
await handle_sse_mcp(mock_scope, mock_receive, mock_send)
|
||||
|
||||
assert run.await_args.args[1:3] == (read_stream, write_stream)
|
||||
assert serve.await_args.args[1:3] == (read_stream, write_stream)
|
||||
assert mock_sse.connect_sse.call_args.args[0]["path"] == "/mcp/sse"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -575,7 +575,7 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client):
|
|||
"""Test getting guardrail info from DB"""
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
|
||||
response = await get_guardrail_info("test-db-guardrail")
|
||||
response: Final = await get_guardrail_info("test-db-guardrail")
|
||||
|
||||
assert response.guardrail_id == "test-db-guardrail"
|
||||
assert response.guardrail_name == "Test DB Guardrail"
|
||||
|
|
@ -584,6 +584,21 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_info_tolerates_invalid_stored_stream_scope(mocker, mock_prisma_client):
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
mock_prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(
|
||||
return_value={
|
||||
**MOCK_DB_GUARDRAIL,
|
||||
"litellm_params": {
|
||||
**MOCK_DB_GUARDRAIL["litellm_params"],
|
||||
"stream_scope": "sometimes",
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
response = await get_guardrail_info("test-db-guardrail")
|
||||
|
||||
assert response.litellm_params.stream_scope is None
|
||||
async def test_get_guardrail_info_normalizes_invalid_scope_from_db(
|
||||
mocker, mock_guardrail_registry, mock_in_memory_handler
|
||||
):
|
||||
|
|
@ -745,6 +760,40 @@ def test_get_guardrails_list_response_includes_guardrail_id():
|
|||
assert response.guardrails[0].guardrail_id == "stable-config-id"
|
||||
|
||||
|
||||
def test_get_guardrails_list_response_tolerates_invalid_config_stream_scope():
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import (
|
||||
_get_guardrails_list_response,
|
||||
)
|
||||
|
||||
response = _get_guardrails_list_response(
|
||||
[
|
||||
{
|
||||
"guardrail_id": "invalid-scope",
|
||||
"guardrail_name": "invalid-scope",
|
||||
"litellm_params": {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"stream_scope": "sometimes",
|
||||
},
|
||||
},
|
||||
{
|
||||
"guardrail_id": "valid-scope",
|
||||
"guardrail_name": "valid-scope",
|
||||
"litellm_params": {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"stream_scope": "STREAMING",
|
||||
},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
assert response.guardrails[0].litellm_params is not None
|
||||
assert response.guardrails[0].litellm_params.stream_scope is None
|
||||
assert response.guardrails[1].litellm_params is not None
|
||||
assert response.guardrails[1].litellm_params.stream_scope == "streaming"
|
||||
|
||||
|
||||
def test_get_provider_specific_params():
|
||||
"""Test getting provider-specific parameters"""
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import _get_fields_from_model
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import json
|
||||
from collections.abc import Iterable, Iterator
|
||||
from typing import ClassVar, Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
|
@ -160,6 +160,43 @@ def test_duplicate_config_guardrail_names_get_distinct_stable_ids():
|
|||
registry_module.guardrail_initializer_registry.pop("dup_name_test", None)
|
||||
|
||||
|
||||
def test_initialize_guardrail_treats_invalid_stored_scope_as_both():
|
||||
from litellm.proxy.guardrails import guardrail_registry as registry_module
|
||||
|
||||
guardrail_type: Final = "invalid_stored_scope_test"
|
||||
|
||||
def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail:
|
||||
return CustomGuardrail(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
event_hook=GuardrailEventHooks(litellm_params.mode),
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
registry_module.guardrail_initializer_registry[guardrail_type] = _initializer
|
||||
try:
|
||||
handler: Final = InMemoryGuardrailHandler()
|
||||
guardrail: Final = Guardrail(
|
||||
guardrail_id="invalid-stored-scope",
|
||||
guardrail_name="invalid-stored-scope",
|
||||
litellm_params={
|
||||
"guardrail": guardrail_type,
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"stream_scope": "sometimes",
|
||||
},
|
||||
)
|
||||
|
||||
parsed_guardrail: Final = handler.initialize_guardrail(guardrail=guardrail, source="db")
|
||||
callback: Final = handler.guardrail_id_to_custom_guardrail["invalid-stored-scope"]
|
||||
|
||||
assert parsed_guardrail["litellm_params"].stream_scope is None
|
||||
assert callback is not None
|
||||
assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True
|
||||
assert callback.should_run_guardrail(data={"stream": True}, event_type=GuardrailEventHooks.pre_call) is True
|
||||
finally:
|
||||
registry_module.guardrail_initializer_registry.pop(guardrail_type, None)
|
||||
|
||||
|
||||
def _register_mode_following_initializer(guardrail_type: str):
|
||||
"""Registers like the shipped initializers do: construct, then add the instance to litellm's callbacks."""
|
||||
import litellm
|
||||
|
|
@ -1605,6 +1642,35 @@ def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance():
|
|||
cb_list[:] = snapshot
|
||||
|
||||
|
||||
def test_configure_callback_scoping_copies_stream_scope_when_constructor_omits_it():
|
||||
from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping
|
||||
|
||||
class _CtorWithoutStreamScope(CustomGuardrail):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(
|
||||
guardrail_name="scoped",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
instance = _CtorWithoutStreamScope()
|
||||
params = LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="streaming")
|
||||
_configure_callback_scoping(instance, "scoped", params)
|
||||
|
||||
assert instance.stream_scope_default == "streaming"
|
||||
assert instance.should_run_guardrail({"stream": True}, GuardrailEventHooks.post_call) is True
|
||||
assert instance.should_run_guardrail({"stream": False}, GuardrailEventHooks.post_call) is False
|
||||
|
||||
|
||||
def test_configure_callback_scoping_tolerates_a_custom_logger_callback():
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy.guardrails.guardrail_registry import _configure_callback_scoping
|
||||
|
||||
callback: Final = CustomLogger()
|
||||
_configure_callback_scoping(callback, "logger-backed", LitellmParams(guardrail="custom", mode="pre_call")) # pyright: ignore[reportArgumentType] # module-path guardrails may be plain CustomLogger
|
||||
assert "stream_scope_by_hook" not in vars(callback)
|
||||
|
||||
|
||||
_ENCRYPTED_PREFIX = "litellm_enc::"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from starlette.datastructures import FormData
|
|||
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from tests._master_key import MASTER_KEY as SHARED_MASTER_KEY
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
|
@ -47,6 +48,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
get_vertex_base_url,
|
||||
is_azure_ai_search_service_level_index_create,
|
||||
gigachat_proxy_route,
|
||||
handle_bedrock_passthrough_router_model,
|
||||
llm_passthrough_factory_proxy_route,
|
||||
milvus_proxy_route,
|
||||
mistral_proxy_route,
|
||||
|
|
@ -61,9 +63,32 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
|
||||
|
||||
|
||||
def _assert_bedrock_processing_data_classification(data: dict[str, object], is_streaming: bool) -> None:
|
||||
streaming_guardrail: Final = CustomGuardrail(
|
||||
guardrail_name="streaming-only",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="streaming",
|
||||
)
|
||||
non_streaming_guardrail: Final = CustomGuardrail(
|
||||
guardrail_name="non-streaming-only",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="non_streaming",
|
||||
)
|
||||
provider_body: Final = data["data"]
|
||||
|
||||
assert streaming_guardrail.should_run_guardrail(data, GuardrailEventHooks.pre_call) is is_streaming
|
||||
assert non_streaming_guardrail.should_run_guardrail(data, GuardrailEventHooks.pre_call) is not is_streaming
|
||||
assert isinstance(provider_body, dict)
|
||||
assert "is_streaming_request" not in provider_body
|
||||
assert "litellm_server_streaming_classification" not in provider_body
|
||||
|
||||
|
||||
class TestVertexPassthroughGetVertexBaseUrl:
|
||||
"""Module-local get_vertex_base_url (trailing slash); rules match common_utils."""
|
||||
|
||||
|
|
@ -413,9 +438,7 @@ class TestVertexAIPassThroughHandler:
|
|||
|
||||
# Mock the vertex handler for global location
|
||||
mock_handler = Mock()
|
||||
mock_handler.get_default_base_target_url.return_value = (
|
||||
"https://aiplatform.googleapis.com/"
|
||||
)
|
||||
mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com/"
|
||||
mock_get_handler.return_value = mock_handler
|
||||
|
||||
# Mock create_pass_through_route to return a function that returns a mock response
|
||||
|
|
@ -1237,9 +1260,7 @@ class TestVertexAIDiscoveryPassThroughHandler:
|
|||
|
||||
# Mock the discovery handler
|
||||
mock_handler = Mock()
|
||||
mock_handler.get_default_base_target_url.return_value = (
|
||||
"https://discoveryengine.googleapis.com"
|
||||
)
|
||||
mock_handler.get_default_base_target_url.return_value = "https://discoveryengine.googleapis.com"
|
||||
mock_get_handler.return_value = mock_handler
|
||||
|
||||
# Mock create_pass_through_route to return a function that returns a mock response
|
||||
|
|
@ -1459,6 +1480,170 @@ class TestBedrockLLMProxyRoute:
|
|||
assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
assert result == "success"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"action, is_streaming",
|
||||
[
|
||||
("converse-stream", True),
|
||||
("invoke-with-response-stream", True),
|
||||
("converse", False),
|
||||
("invoke", False),
|
||||
],
|
||||
)
|
||||
async def test_bedrock_direct_actions_classify_guardrail_stream_scope(
|
||||
self, action: str, is_streaming: bool
|
||||
) -> None:
|
||||
mock_request: Final = Mock()
|
||||
mock_request.method = "POST"
|
||||
mock_processor: Final = Mock()
|
||||
mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success")
|
||||
request_body: Final = {"messages": [{"role": "user", "content": "test"}]}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body",
|
||||
return_value=request_body,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
) as processor_constructor,
|
||||
):
|
||||
result: Final = await bedrock_llm_proxy_route(
|
||||
endpoint=f"model/test-model/{action}",
|
||||
request=mock_request,
|
||||
fastapi_response=Mock(),
|
||||
user_api_key_dict=Mock(),
|
||||
)
|
||||
|
||||
assert result == "success"
|
||||
processing_data: Final = processor_constructor.call_args.kwargs["data"]
|
||||
_assert_bedrock_processing_data_classification(processing_data, is_streaming)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"action, is_streaming",
|
||||
[
|
||||
("converse-stream", True),
|
||||
("invoke-with-response-stream", True),
|
||||
("converse", False),
|
||||
("invoke", False),
|
||||
],
|
||||
)
|
||||
async def test_bedrock_router_actions_classify_guardrail_stream_scope(
|
||||
self, action: str, is_streaming: bool
|
||||
) -> None:
|
||||
mock_request: Final = Mock()
|
||||
mock_request.method = "POST"
|
||||
mock_processor: Final = Mock()
|
||||
mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success")
|
||||
request_body: Final = {"messages": [{"role": "user", "content": "test"}]}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
) as processor_constructor:
|
||||
result: Final = await handle_bedrock_passthrough_router_model(
|
||||
model="test-model",
|
||||
endpoint=f"model/test-model/{action}",
|
||||
request=mock_request,
|
||||
request_body=request_body,
|
||||
llm_router=Mock(),
|
||||
user_api_key_dict=Mock(),
|
||||
proxy_logging_obj=Mock(),
|
||||
general_settings={},
|
||||
proxy_config=None,
|
||||
select_data_generator=None,
|
||||
user_model=None,
|
||||
user_temperature=None,
|
||||
user_request_timeout=None,
|
||||
user_max_tokens=None,
|
||||
user_api_base=None,
|
||||
version=None,
|
||||
)
|
||||
|
||||
assert result == "success"
|
||||
processing_data: Final = processor_constructor.call_args.kwargs["data"]
|
||||
_assert_bedrock_processing_data_classification(processing_data, is_streaming)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("model_id", "action"),
|
||||
[
|
||||
("my-converse-stream-model", "converse"),
|
||||
("my-invoke-with-response-stream-model", "invoke"),
|
||||
],
|
||||
)
|
||||
async def test_bedrock_direct_model_id_does_not_imply_streaming(self, model_id: str, action: str) -> None:
|
||||
mock_request: Final = Mock()
|
||||
mock_request.method = "POST"
|
||||
mock_processor: Final = Mock()
|
||||
mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success")
|
||||
request_body: Final = {"messages": [{"role": "user", "content": "test"}]}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._read_request_body",
|
||||
return_value=request_body,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
) as processor_constructor,
|
||||
):
|
||||
result: Final = await bedrock_llm_proxy_route(
|
||||
endpoint=f"/model/{model_id}/{action}",
|
||||
request=mock_request,
|
||||
fastapi_response=Mock(),
|
||||
user_api_key_dict=Mock(),
|
||||
)
|
||||
|
||||
assert result == "success"
|
||||
processing_data: Final = processor_constructor.call_args.kwargs["data"]
|
||||
_assert_bedrock_processing_data_classification(processing_data, is_streaming=False)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("model_id", "action"),
|
||||
[
|
||||
("my-converse-stream-model", "converse"),
|
||||
("my-invoke-with-response-stream-model", "invoke"),
|
||||
],
|
||||
)
|
||||
async def test_bedrock_router_model_id_does_not_imply_streaming(self, model_id: str, action: str) -> None:
|
||||
mock_request: Final = Mock()
|
||||
mock_request.method = "POST"
|
||||
mock_processor: Final = Mock()
|
||||
mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success")
|
||||
request_body: Final = {"messages": [{"role": "user", "content": "test"}]}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
) as processor_constructor:
|
||||
result: Final = await handle_bedrock_passthrough_router_model(
|
||||
model=model_id,
|
||||
endpoint=f"/model/{model_id}/{action}",
|
||||
request=mock_request,
|
||||
request_body=request_body,
|
||||
llm_router=Mock(),
|
||||
user_api_key_dict=Mock(),
|
||||
proxy_logging_obj=Mock(),
|
||||
general_settings={},
|
||||
proxy_config=None,
|
||||
select_data_generator=None,
|
||||
user_model=None,
|
||||
user_temperature=None,
|
||||
user_request_timeout=None,
|
||||
user_max_tokens=None,
|
||||
user_api_base=None,
|
||||
version=None,
|
||||
)
|
||||
|
||||
assert result == "success"
|
||||
processing_data: Final = processor_constructor.call_args.kwargs["data"]
|
||||
_assert_bedrock_processing_data_classification(processing_data, is_streaming=False)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_error_handling_returns_actual_error(self):
|
||||
"""
|
||||
|
|
@ -1879,7 +2064,6 @@ class TestBedrockAgentRuntimePassthroughToggle:
|
|||
|
||||
|
||||
class TestBedrockAgentRuntimePassthroughVirtualKeyLeak:
|
||||
|
||||
VKEY: Final = "sk-litellm-victim-key"
|
||||
MASTER_KEY: Final = SHARED_MASTER_KEY
|
||||
ENDPOINT: Final = "knowledgebases/KB1234567/retrieve"
|
||||
|
|
@ -2075,10 +2259,10 @@ class TestVLLMProxyRoute:
|
|||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
|
||||
return_value=True,
|
||||
)
|
||||
@patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation
|
||||
async def test_vllm_proxy_route_with_router_model(
|
||||
self, mock_llm_router, mock_is_router, mock_get_body
|
||||
):
|
||||
@patch(
|
||||
"litellm.proxy.proxy_server.llm_router"
|
||||
) # test-quality-ok: patching litellm internal for unit test isolation
|
||||
async def test_vllm_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body):
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
|
|
@ -2111,9 +2295,7 @@ class TestVLLMProxyRoute:
|
|||
@patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.llm_passthrough_factory_proxy_route"
|
||||
)
|
||||
async def test_vllm_proxy_route_fallback_to_factory(
|
||||
self, mock_factory_route, mock_is_router, mock_get_body
|
||||
):
|
||||
async def test_vllm_proxy_route_fallback_to_factory(self, mock_factory_route, mock_is_router, mock_get_body):
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_fastapi_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock()
|
||||
|
|
@ -2140,10 +2322,10 @@ class TestGigachatProxyRoute:
|
|||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
|
||||
return_value=True,
|
||||
)
|
||||
@patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation
|
||||
async def test_gigachat_proxy_route_with_router_model(
|
||||
self, mock_llm_router, mock_is_router, mock_get_body
|
||||
):
|
||||
@patch(
|
||||
"litellm.proxy.proxy_server.llm_router"
|
||||
) # test-quality-ok: patching litellm internal for unit test isolation
|
||||
async def test_gigachat_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body):
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.headers = {"content-type": "application/json"}
|
||||
|
|
@ -2396,21 +2578,25 @@ class TestGigachatProxyRoute:
|
|||
|
||||
return _inner()
|
||||
|
||||
with patch.object(
|
||||
processor,
|
||||
"common_processing_pre_call_logic",
|
||||
new=AsyncMock(
|
||||
return_value=(
|
||||
processor.data,
|
||||
processor.data["litellm_logging_obj"],
|
||||
)
|
||||
with (
|
||||
patch.object(
|
||||
processor,
|
||||
"common_processing_pre_call_logic",
|
||||
new=AsyncMock(
|
||||
return_value=(
|
||||
processor.data,
|
||||
processor.data["litellm_logging_obj"],
|
||||
)
|
||||
),
|
||||
),
|
||||
patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.common_request_processing.route_request",
|
||||
new=_fake_route_request,
|
||||
),
|
||||
patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers",
|
||||
return_value={"x-litellm-call-id": "call-123"},
|
||||
),
|
||||
), patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.common_request_processing.route_request",
|
||||
new=_fake_route_request,
|
||||
), patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers",
|
||||
return_value={"x-litellm-call-id": "call-123"},
|
||||
):
|
||||
result = await processor.base_passthrough_process_llm_request(
|
||||
request=mock_request,
|
||||
|
|
@ -4223,7 +4409,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
|
|||
@pytest.mark.parametrize(
|
||||
("credential", "authenticated"),
|
||||
[
|
||||
pytest.param("modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential"),
|
||||
pytest.param(
|
||||
"modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential"
|
||||
),
|
||||
pytest.param(
|
||||
LITELLM_JWT,
|
||||
UserAPIKeyAuth(api_key=LITELLM_JWT, user_id="jwt-subject"),
|
||||
|
|
@ -4298,7 +4486,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
|
|||
(b"x-goog-api-key", b"AIza-real-google-api-key"),
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
authenticated=UserAPIKeyAuth(
|
||||
api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
),
|
||||
)
|
||||
assert raised is None
|
||||
assert forwarded is not None
|
||||
|
|
@ -4440,10 +4630,14 @@ class TestAnthropicPassthroughVirtualKeyLeak:
|
|||
raised, forwarded = await self._run(
|
||||
monkeypatch,
|
||||
[(header, value), (b"anthropic-version", b"2023-06-01"), (b"content-type", b"application/json")],
|
||||
authenticated=UserAPIKeyAuth(api_key="sk-ant-api03-callers-own-key", user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
authenticated=UserAPIKeyAuth(
|
||||
api_key="sk-ant-api03-callers-own-key", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
),
|
||||
master_key=None,
|
||||
)
|
||||
assert raised is None, "with no master key the proxy authenticated nothing, so nothing of the caller's is a LiteLLM secret"
|
||||
assert raised is None, (
|
||||
"with no master key the proxy authenticated nothing, so nothing of the caller's is a LiteLLM secret"
|
||||
)
|
||||
assert forwarded is not None
|
||||
assert forwarded.get(header.decode()) == value.decode()
|
||||
|
||||
|
|
@ -5231,7 +5425,9 @@ class TestTranscribeProxyRoute:
|
|||
) -> None:
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route = upstream.post(TRANSCRIBE_UPSTREAM)
|
||||
response = transcribe_client.post("/transcribe/StartTranscriptionJob", json={**dict(self.START_JOB_BODY), **body})
|
||||
response = transcribe_client.post(
|
||||
"/transcribe/StartTranscriptionJob", json={**dict(self.START_JOB_BODY), **body}
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert member in response.json()["detail"]
|
||||
|
|
@ -6873,9 +7069,7 @@ class TestAzureBodyModelGroupRelay:
|
|||
AZURE_SPEECH_SHORT_AUDIO_ENDPOINT: Final = "/speech/recognition/conversation/cognitiveservices/v1"
|
||||
AZURE_SPEECH_BATCH_ENDPOINT: Final = "/speechtotext/v3.2/transcriptions"
|
||||
AZURE_SPEECH_FAST_ENDPOINT: Final = "/speechtotext/transcriptions:transcribe"
|
||||
AZURE_SPEECH_PCM16_HEADER: Final = (
|
||||
b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00"
|
||||
)
|
||||
AZURE_SPEECH_PCM16_HEADER: Final = b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00"
|
||||
AZURE_SPEECH_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + b"\x00" * 3072
|
||||
AZURE_SPEECH_WAV_SECONDS: Final = 3072 / (16000 * 2)
|
||||
AZURE_SPEECH_NON_UTF8_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + bytes(range(256)) * 12
|
||||
|
|
@ -6912,9 +7106,9 @@ class TestAzureSpeechProxyRoute:
|
|||
|
||||
def test_short_audio_forwards_raw_wav_bytes_with_server_key(self, azure_speech_client: TestClient) -> None:
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
route = upstream.post(
|
||||
f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}"
|
||||
).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT))
|
||||
route = upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock(
|
||||
return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)
|
||||
)
|
||||
|
||||
response = azure_speech_client.post(
|
||||
f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}",
|
||||
|
|
@ -7390,9 +7584,9 @@ class TestAzureSpeechRawBodyThroughRealAuth:
|
|||
self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, body: bytes
|
||||
) -> None:
|
||||
with respx.mock(assert_all_called=True) as upstream, caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"):
|
||||
route = upstream.post(
|
||||
f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}"
|
||||
).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT))
|
||||
route = upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock(
|
||||
return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)
|
||||
)
|
||||
|
||||
response = self._post_wav(
|
||||
monkeypatch, f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", "sk-master-key", body=body
|
||||
|
|
@ -7416,9 +7610,9 @@ class TestAzureSpeechRawBodyThroughRealAuth:
|
|||
) -> None:
|
||||
boundary: Final = "lit7939boundary"
|
||||
multipart_body: Final = (
|
||||
f"--{boundary}\r\nContent-Disposition: form-data; name=\"definition\"\r\n\r\n".encode()
|
||||
f'--{boundary}\r\nContent-Disposition: form-data; name="definition"\r\n\r\n'.encode()
|
||||
+ json.dumps({"locales": ["en-US"]}).encode()
|
||||
+ f"\r\n--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"eagle.wav\"\r\n"
|
||||
+ f'\r\n--{boundary}\r\nContent-Disposition: form-data; name="audio"; filename="eagle.wav"\r\n'
|
||||
"Content-Type: audio/wav\r\n\r\n".encode()
|
||||
+ AZURE_SPEECH_NON_UTF8_WAV_BYTES
|
||||
+ f"\r\n--{boundary}--\r\n".encode()
|
||||
|
|
@ -8170,9 +8364,7 @@ class TestTinyFishProxyRoute:
|
|||
assert response.status_code == 200
|
||||
assert json.loads(route.calls.last.request.content)["use_vault"] is True
|
||||
|
||||
def test_returns_401_on_missing_api_key(
|
||||
self, tinyfish_client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def test_returns_401_on_missing_api_key(self, tinyfish_client: TestClient, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("TINYFISH_API_KEY")
|
||||
|
||||
with respx.mock:
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ import logging
|
|||
import os
|
||||
import sys
|
||||
import zlib
|
||||
from collections.abc import Callable, Mapping
|
||||
from collections.abc import AsyncIterator, Callable, Mapping
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from dataclasses import dataclass
|
||||
from io import BytesIO
|
||||
|
|
@ -27,6 +27,7 @@ from starlette.datastructures import UploadFile as StarletteUploadFile
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features
|
||||
|
|
@ -55,6 +56,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
|
|||
)
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
from litellm.types import utils as types_utils
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
|
|
@ -1638,6 +1640,223 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream():
|
|||
assert logging_obj.model_call_details["stream"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("body_stream", [None, True], ids=["stream-absent", "stream-true"])
|
||||
async def test_pass_through_request_preserves_caller_streaming_request_field(body_stream):
|
||||
captured_hook_data: dict[str, object] = {}
|
||||
|
||||
async def capture_pre_call(user_api_key_dict, data, call_type, endpoint_type: EndpointType):
|
||||
captured_hook_data.update(data)
|
||||
return data
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor"
|
||||
) as mock_chunk_processor:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=capture_pre_call)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
|
||||
upstream_response = MagicMock()
|
||||
upstream_response.status_code = 200
|
||||
upstream_response.headers = {}
|
||||
upstream_response.raise_for_status = MagicMock()
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
async def _empty_chunks(*args, **kwargs):
|
||||
return
|
||||
yield # pragma: no cover
|
||||
|
||||
mock_chunk_processor.return_value = _empty_chunks()
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = httpx.URL(
|
||||
"http://test-proxy.com/gemini/v1beta/models/gemini-pro:streamGenerateContent"
|
||||
)
|
||||
mock_request.scope = {"path": "/gemini/v1beta/models/gemini-pro:streamGenerateContent"}
|
||||
request_body: Final = {
|
||||
"contents": [{"parts": [{"text": "hi"}]}],
|
||||
"is_streaming_request": "caller-value",
|
||||
**({"stream": True} if body_stream is True else {}),
|
||||
}
|
||||
mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode())
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:streamGenerateContent",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert captured_hook_data.get("is_streaming_request") == "caller-value"
|
||||
assert captured_hook_data.get("stream") is body_stream
|
||||
|
||||
upstream_json = async_client.build_request.call_args.kwargs["json"]
|
||||
assert upstream_json["is_streaming_request"] == "caller-value"
|
||||
assert "litellm_server_streaming_classification" not in upstream_json
|
||||
assert upstream_json["contents"] == request_body["contents"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_pass_through_drops_marker_after_hook_rebuilds_body_from_json():
|
||||
async def json_rebuilding_pre_call(user_api_key_dict, data, call_type, endpoint_type: EndpointType):
|
||||
rebuilt = json.loads(json.dumps({k: v for k, v in data.items() if k != "litellm_logging_obj"}))
|
||||
return {**rebuilt, "litellm_logging_obj": data["litellm_logging_obj"]}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor"
|
||||
) as mock_chunk_processor:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=json_rebuilding_pre_call)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
|
||||
upstream_response = MagicMock()
|
||||
upstream_response.status_code = 200
|
||||
upstream_response.headers = {}
|
||||
upstream_response.raise_for_status = MagicMock()
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
async def _empty_chunks(*args, **kwargs):
|
||||
return
|
||||
yield # pragma: no cover
|
||||
|
||||
mock_chunk_processor.return_value = _empty_chunks()
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/openai/v1/chat/completions")
|
||||
mock_request.scope = {"path": "/openai/v1/chat/completions"}
|
||||
request_body: Final = {
|
||||
"model": "gpt-5-mini",
|
||||
"stream": True,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode())
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="https://api.openai.com/v1/chat/completions",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
upstream_json = async_client.build_request.call_args.kwargs["json"]
|
||||
assert upstream_json == request_body, upstream_json
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("route_stream", "body_stream", "expected_streaming"),
|
||||
[
|
||||
(True, False, False),
|
||||
(False, True, True),
|
||||
(True, None, True),
|
||||
(None, None, False),
|
||||
],
|
||||
ids=["body-disables-route-stream", "body-enables-streaming", "route-enables-absent-body", "both-absent"],
|
||||
)
|
||||
async def test_passthrough_guardrails_follow_effective_relay_stream_decision(
|
||||
route_stream: bool | None,
|
||||
body_stream: bool | None,
|
||||
expected_streaming: bool,
|
||||
):
|
||||
async def return_pre_call_data(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
data: dict[str, object],
|
||||
call_type: str,
|
||||
endpoint_type: EndpointType,
|
||||
) -> dict[str, object]:
|
||||
return dict(data)
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor"
|
||||
) as mock_chunk_processor:
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=return_pre_call_data)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
|
||||
upstream_response = MagicMock()
|
||||
upstream_response.status_code = 200
|
||||
upstream_response.headers = {}
|
||||
upstream_response.aread = AsyncMock(return_value=b"{}")
|
||||
upstream_response.text = "{}"
|
||||
upstream_response.raise_for_status = MagicMock()
|
||||
|
||||
async_client = MagicMock()
|
||||
async_client.build_request = MagicMock(return_value=MagicMock())
|
||||
async_client.send = AsyncMock(return_value=upstream_response)
|
||||
async_client.request = AsyncMock(return_value=upstream_response)
|
||||
mock_get_client.return_value = MagicMock(client=async_client)
|
||||
|
||||
async def empty_chunks() -> AsyncIterator[bytes]:
|
||||
yield b""
|
||||
|
||||
mock_chunk_processor.return_value = empty_chunks()
|
||||
|
||||
request_body: Final = {
|
||||
"message": "hello",
|
||||
**({"stream": body_stream} if body_stream is not None else {}),
|
||||
}
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = httpx.URL("http://test-proxy.com/guardrail-stream-scope")
|
||||
mock_request.scope = {"path": "/guardrail-stream-scope"}
|
||||
mock_request.body = AsyncMock(return_value=json.dumps(request_body).encode())
|
||||
mock_request.headers = Headers({"content-type": "application/json"})
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://upstream.test/guardrail-stream-scope",
|
||||
custom_headers={},
|
||||
user_api_key_dict=MagicMock(),
|
||||
stream=route_stream,
|
||||
)
|
||||
hook_data: Final[dict[str, object]] = mock_proxy_logging.pre_call_hook.call_args.kwargs["data"]
|
||||
|
||||
streaming_guardrail: Final = CustomGuardrail(
|
||||
guardrail_name="streaming-only",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="streaming",
|
||||
)
|
||||
non_streaming_guardrail: Final = CustomGuardrail(
|
||||
guardrail_name="non-streaming-only",
|
||||
default_on=True,
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
stream_scope="non_streaming",
|
||||
)
|
||||
assert streaming_guardrail.should_run_guardrail(hook_data, GuardrailEventHooks.pre_call) is expected_streaming
|
||||
assert (
|
||||
non_streaming_guardrail.should_run_guardrail(hook_data, GuardrailEventHooks.pre_call) is not expected_streaming
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_sse_response_marks_logging_obj_as_stream():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -96,6 +96,20 @@ class AlwaysPassGuardrail(CustomGuardrail):
|
|||
return None
|
||||
|
||||
|
||||
class StreamScopedPassGuardrail(CustomGuardrail):
|
||||
def __init__(self, guardrail_name: str, stream_scope: object):
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
stream_scope=stream_scope,
|
||||
)
|
||||
self.calls = 0
|
||||
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.calls += 1
|
||||
|
||||
|
||||
class PassthroughBlockGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that blocks using the legacy passthrough contract."""
|
||||
|
||||
|
|
@ -378,6 +392,46 @@ async def test_block_carries_original_guardrail_exception(monkeypatch):
|
|||
assert result.original_exception.detail == "Content policy violation"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pipeline_step_honors_stream_scope(monkeypatch):
|
||||
stream_only = StreamScopedPassGuardrail(guardrail_name="stream-only", stream_scope="streaming")
|
||||
later = AlwaysFailGuardrail(guardrail_name="later-block")
|
||||
monkeypatch.setattr(litellm, "callbacks", [stream_only, later])
|
||||
steps = [
|
||||
PipelineStep(guardrail="stream-only", on_fail="block", on_pass="allow"),
|
||||
PipelineStep(guardrail="later-block", on_fail="block", on_pass="allow"),
|
||||
]
|
||||
|
||||
skipped = await PipelineExecutor.execute_steps(
|
||||
steps=steps,
|
||||
mode="pre_call",
|
||||
data={"messages": [{"role": "user", "content": "hi"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="stream-scope",
|
||||
)
|
||||
assert stream_only.calls == 0
|
||||
assert later.calls == 1
|
||||
assert skipped.step_results[0].outcome == "skip"
|
||||
assert skipped.step_results[0].action_taken == "next"
|
||||
assert skipped.terminal_action == "block"
|
||||
|
||||
stream_only.calls = 0
|
||||
later.calls = 0
|
||||
ran = await PipelineExecutor.execute_steps(
|
||||
steps=steps,
|
||||
mode="pre_call",
|
||||
data={"messages": [{"role": "user", "content": "hi"}], "stream": True},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="stream-scope",
|
||||
)
|
||||
assert stream_only.calls == 1
|
||||
assert later.calls == 0
|
||||
assert ran.step_results[0].outcome == "pass"
|
||||
assert ran.terminal_action == "allow"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unsupported_mode_yields_error_outcome_without_exception(monkeypatch):
|
||||
"""An unexpected hook mode must surface as an error outcome (carrying no
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from pydantic import ValidationError as PydanticValidationError
|
|||
from starlette.datastructures import Headers
|
||||
|
||||
import litellm
|
||||
from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER
|
||||
from litellm.proxy._types import AddTeamCallback, ProxyException, TeamCallbackMetadata, UserAPIKeyAuth
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
KeyAndTeamLoggingSettings,
|
||||
|
|
@ -850,8 +851,13 @@ def test_initial_snapshot_refresh_clears_a_previous_guardrail_checkpoint() -> No
|
|||
from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot
|
||||
|
||||
logging_obj: Final = Logging(
|
||||
model="test-model", messages=[], stream=False, call_type="acompletion",
|
||||
start_time=datetime.now(), litellm_call_id="new-request", function_id="new-request",
|
||||
model="test-model",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="new-request",
|
||||
function_id="new-request",
|
||||
)
|
||||
logging_obj.shadow_eval_request_snapshot = GuardrailRequestSnapshot.capture(
|
||||
{"messages": [{"role": "user", "content": "previous request"}]},
|
||||
|
|
@ -913,8 +919,14 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l
|
|||
}
|
||||
data: Final = {"messages": messages, "metadata": metadata, "proxy_server_request": {}}
|
||||
logging_obj: Final = Logging(
|
||||
model="test-model", messages=messages, stream=False, call_type="acompletion",
|
||||
start_time=datetime.now(), litellm_call_id="mask-spend", function_id="mask-spend", kwargs=data,
|
||||
model="test-model",
|
||||
messages=messages,
|
||||
stream=False,
|
||||
call_type="acompletion",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="mask-spend",
|
||||
function_id="mask-spend",
|
||||
kwargs=data,
|
||||
)
|
||||
data["litellm_logging_obj"] = logging_obj
|
||||
refresh_proxy_server_request_body_snapshot(data, guardrails_applied=True)
|
||||
|
|
@ -928,9 +940,13 @@ async def test_post_guardrail_snapshot_preserves_logging_only_masking_in_spend_l
|
|||
kwargs, _ = await guardrail.async_logging_hook(
|
||||
kwargs=logging_obj.model_call_details, result=None, call_type="acompletion"
|
||||
)
|
||||
stored: Final = json.loads(_get_proxy_server_request_for_spend_logs_payload(
|
||||
metadata={}, litellm_params=kwargs["litellm_params"], kwargs=kwargs,
|
||||
))
|
||||
stored: Final = json.loads(
|
||||
_get_proxy_server_request_for_spend_logs_payload(
|
||||
metadata={},
|
||||
litellm_params=kwargs["litellm_params"],
|
||||
kwargs=kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
assert kwargs["messages"] == [{"role": "user", "content": "email [EMAIL]"}]
|
||||
assert stored["messages"] == kwargs["messages"]
|
||||
|
|
@ -1089,6 +1105,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields():
|
|||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"mock_response": "free response",
|
||||
"mock_tool_calls": [{"id": "call_1"}],
|
||||
"is_streaming_request": "caller-value",
|
||||
"disable_global_guardrails": True,
|
||||
"enable_prompt_caching": True,
|
||||
"routing_decision": {"cause": "forged", "routed_model": "spoofed"},
|
||||
|
|
@ -1114,6 +1131,7 @@ async def test_add_litellm_data_to_request_strips_user_control_fields():
|
|||
assert "enable_prompt_caching" not in updated
|
||||
assert "routing_decision" not in updated
|
||||
assert "litellm_gateway_injected_cache" not in updated
|
||||
assert updated["is_streaming_request"] == "caller-value"
|
||||
assert "weights" not in updated
|
||||
assert "_router_weights" not in updated
|
||||
assert "weights" not in updated["proxy_server_request"]["body"]
|
||||
|
|
@ -8803,20 +8821,27 @@ async def test_mcp_credentials_only_removed_from_logging_copies(path: str, custo
|
|||
request.headers = Headers(request.headers)
|
||||
settings: Final = {"mcp_client_side_auth_header_name": custom_auth, "user_header_name": "x-user-id"}
|
||||
server: Final = MCPServer(
|
||||
server_id="header-test", name="header-test", transport="http", url="https://example.com/mcp",
|
||||
server_id="header-test",
|
||||
name="header-test",
|
||||
transport="http",
|
||||
url="https://example.com/mcp",
|
||||
extra_headers=["x-service-token", "x-user-id"],
|
||||
)
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.general_settings", settings),
|
||||
patch.dict(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers",
|
||||
{"header-test": server}, clear=True,
|
||||
{"header-test": server},
|
||||
clear=True,
|
||||
),
|
||||
):
|
||||
updated: Final = await add_litellm_data_to_request(
|
||||
data={"model": "test-model", "messages": [{"role": "user", "content": "hello"}]},
|
||||
request=request, user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
proxy_config=MagicMock(), general_settings=settings, version="test",
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings=settings,
|
||||
version="test",
|
||||
)
|
||||
for header_dict in _all_header_dicts(updated, metadata_name):
|
||||
assert not any(value in json.dumps(header_dict) for value in secrets.values())
|
||||
|
|
@ -8836,7 +8861,10 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback():
|
|||
data=AddTeamCallback(
|
||||
callback_name="signoz",
|
||||
callback_type="success",
|
||||
callback_vars={"signoz_ingestion_key": "team-key", "signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443"},
|
||||
callback_vars={
|
||||
"signoz_ingestion_key": "team-key",
|
||||
"signoz_ingestion_endpoint": "https://ingest.eu.signoz.cloud:443",
|
||||
},
|
||||
),
|
||||
team_callback_settings_obj=None,
|
||||
)
|
||||
|
|
@ -8856,6 +8884,54 @@ def test_signoz_callback_vars_are_scoped_to_the_signoz_callback():
|
|||
assert under_other.callback_vars == {"langfuse_host": "https://cloud.langfuse.com"}
|
||||
|
||||
|
||||
def test_body_snapshot_excludes_the_server_streaming_marker() -> None:
|
||||
from litellm.constants import SERVER_STREAMING_CLASSIFICATION_KEY, SERVER_STREAMING_CLASSIFICATION_MARKER
|
||||
from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot
|
||||
|
||||
proxy_request: Final = {"body": {}}
|
||||
data: Final = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: SERVER_STREAMING_CLASSIFICATION_MARKER,
|
||||
"proxy_server_request": proxy_request,
|
||||
}
|
||||
|
||||
refresh_proxy_server_request_body_snapshot(data)
|
||||
|
||||
assert proxy_request == {"body": {"messages": [{"role": "user", "content": "hi"}]}}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"marker",
|
||||
[
|
||||
SERVER_STREAMING_CLASSIFICATION_MARKER,
|
||||
json.loads(json.dumps(SERVER_STREAMING_CLASSIFICATION_MARKER)),
|
||||
],
|
||||
ids=["enum", "json-string"],
|
||||
)
|
||||
def test_body_snapshot_drops_only_the_marker_and_keeps_caller_value(marker: str) -> None:
|
||||
from litellm.proxy.litellm_pre_call_utils import refresh_proxy_server_request_body_snapshot
|
||||
|
||||
marker_request: Final = {"body": {}}
|
||||
marker_data: Final = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: marker,
|
||||
"proxy_server_request": marker_request,
|
||||
}
|
||||
|
||||
refresh_proxy_server_request_body_snapshot(marker_data)
|
||||
|
||||
assert SERVER_STREAMING_CLASSIFICATION_KEY not in marker_request["body"], marker_request
|
||||
|
||||
caller_request: Final = {"body": {}}
|
||||
caller_data: Final = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
SERVER_STREAMING_CLASSIFICATION_KEY: "caller-value",
|
||||
"proxy_server_request": caller_request,
|
||||
}
|
||||
|
||||
refresh_proxy_server_request_body_snapshot(caller_data)
|
||||
|
||||
assert caller_request["body"][SERVER_STREAMING_CLASSIFICATION_KEY] == "caller-value", caller_request
|
||||
def test_arize_otlp_protocol_on_a_key_logging_entry_reaches_the_destination(monkeypatch):
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
from litellm.proxy.litellm_pre_call_utils import resolve_tenant_otel_destinations
|
||||
|
|
|
|||
|
|
@ -26,7 +26,6 @@ from litellm.proxy.litellm_pre_call_utils import (
|
|||
from litellm.types.utils import CustomPricingLiteLLMParams
|
||||
|
||||
|
||||
|
||||
def _make_request_mock() -> Request:
|
||||
request_mock = MagicMock(spec=Request)
|
||||
request_mock.url.path = "/v1/chat/completions"
|
||||
|
|
@ -58,9 +57,7 @@ class TestStripClientPricingOverrides:
|
|||
# The strip set is built from the model so additions are picked up
|
||||
# automatically — this test guards against the model and the strip
|
||||
# set drifting apart if someone replaces the auto-derivation later.
|
||||
assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset(
|
||||
CustomPricingLiteLLMParams.model_fields.keys()
|
||||
)
|
||||
assert _CLIENT_PRICING_CONTROL_FIELDS == frozenset(CustomPricingLiteLLMParams.model_fields.keys())
|
||||
# Sanity: the obvious top-level pricing fields are in the set.
|
||||
for field in (
|
||||
"input_cost_per_token",
|
||||
|
|
@ -184,9 +181,7 @@ class TestStripClientPricingOverrides:
|
|||
verbose_proxy_logger.setLevel(logging.DEBUG)
|
||||
with caplog.at_level(logging.DEBUG, logger=verbose_proxy_logger.name):
|
||||
_strip_client_pricing_overrides({"model": "gpt-4", "temperature": 0.7})
|
||||
assert not any(
|
||||
"pricing" in record.getMessage().lower() for record in caplog.records
|
||||
)
|
||||
assert not any("pricing" in record.getMessage().lower() for record in caplog.records)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -211,6 +206,26 @@ async def test_add_litellm_data_to_request_strips_root_pricing_fields():
|
|||
assert "output_cost_per_token" not in updated
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_preserves_caller_streaming_request():
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"is_streaming_request": True,
|
||||
}
|
||||
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=_make_request_mock(),
|
||||
user_api_key_dict=_user_api_key_auth(),
|
||||
proxy_config=MagicMock(),
|
||||
general_settings={},
|
||||
version="test-version",
|
||||
)
|
||||
|
||||
assert updated["is_streaming_request"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_strips_client_disconnect_metadata():
|
||||
data = {
|
||||
|
|
@ -318,9 +333,7 @@ async def test_add_litellm_data_to_request_skips_strip_with_team_opt_in():
|
|||
"input_cost_per_token": 0.0001,
|
||||
}
|
||||
|
||||
user_auth = _user_api_key_auth(
|
||||
team_metadata={"allow_client_pricing_override": True}
|
||||
)
|
||||
user_auth = _user_api_key_auth(team_metadata={"allow_client_pricing_override": True})
|
||||
updated = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=_make_request_mock(),
|
||||
|
|
|
|||
|
|
@ -2,12 +2,19 @@
|
|||
Test case normalization in LitellmParams for all guardrail types
|
||||
"""
|
||||
|
||||
from typing import Literal
|
||||
import logging
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.types.guardrails import BaseLitellmParams, LitellmParams
|
||||
from litellm.types.guardrails import (
|
||||
BaseLitellmParams,
|
||||
LitellmParams,
|
||||
runtime_stream_scope,
|
||||
stored_stream_scope,
|
||||
with_tolerated_stream_scope,
|
||||
)
|
||||
|
||||
|
||||
class TestLitellmParamsCaseNormalization:
|
||||
|
|
@ -184,3 +191,80 @@ class TestSensitiveDataRoutingValidation:
|
|||
on_sensitive_data="BLOCK",
|
||||
)
|
||||
assert params.on_sensitive_data == "block"
|
||||
|
||||
|
||||
class TestStreamScopeValidation:
|
||||
def test_scalar_is_case_normalized(self):
|
||||
params = LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="Streaming")
|
||||
assert params.stream_scope == "streaming"
|
||||
|
||||
def test_map_keys_and_values_are_normalized(self):
|
||||
params = LitellmParams(
|
||||
guardrail="bedrock",
|
||||
mode=["pre_call", "post_call"],
|
||||
stream_scope={"Pre_Call": "Both", "POST_CALL": "Non_Streaming"},
|
||||
)
|
||||
assert params.stream_scope == {"pre_call": "both", "post_call": "non_streaming"}
|
||||
|
||||
def test_invalid_scalar_is_rejected(self):
|
||||
with pytest.raises(ValidationError, match="stream_scope must be one of"):
|
||||
LitellmParams(guardrail="bedrock", mode="post_call", stream_scope="chunks")
|
||||
|
||||
def test_invalid_map_key_is_rejected(self):
|
||||
with pytest.raises(ValidationError, match="stream_scope keys must be guardrail modes"):
|
||||
LitellmParams(guardrail="bedrock", mode="post_call", stream_scope={"not_a_mode": "both"})
|
||||
|
||||
def test_invalid_map_value_is_rejected(self):
|
||||
with pytest.raises(ValidationError, match="stream_scope must be one of"):
|
||||
LitellmParams(guardrail="bedrock", mode="post_call", stream_scope={"post_call": "sometimes"})
|
||||
|
||||
def test_runtime_stream_scope_normalizes_direct_constructor_maps(self):
|
||||
default, by_hook = runtime_stream_scope({"Pre_Call": "streaming"})
|
||||
assert default == "both"
|
||||
assert dict(by_hook) == {"pre_call": "streaming"}
|
||||
|
||||
def test_runtime_stream_scope_rejects_invalid_direct_input(self):
|
||||
with pytest.raises(ValueError, match="stream_scope must be one of"):
|
||||
runtime_stream_scope("chunks")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"value, expected",
|
||||
[
|
||||
("streaming", "streaming"),
|
||||
("non_streaming", "non_streaming"),
|
||||
("both", "both"),
|
||||
({"pre_call": "streaming"}, {"pre_call": "streaming"}),
|
||||
("sometimes", None),
|
||||
({"pre_call": "sometimes"}, None),
|
||||
],
|
||||
)
|
||||
def test_stored_stream_scope_tolerates_invalid_values(
|
||||
self,
|
||||
value: object,
|
||||
expected: object,
|
||||
caplog: pytest.LogCaptureFixture,
|
||||
) -> None:
|
||||
caplog.set_level(logging.WARNING)
|
||||
|
||||
result: Final = stored_stream_scope(value)
|
||||
|
||||
assert result == expected
|
||||
if expected is None:
|
||||
assert f"Ignoring invalid stored stream_scope value of type {type(value).__name__}" in caplog.text
|
||||
assert "sometimes" not in caplog.text
|
||||
|
||||
def test_tolerated_stream_scope_rewrites_only_the_scope_field(self) -> None:
|
||||
params: Final = {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"stream_scope": "sometimes",
|
||||
}
|
||||
|
||||
tolerated: Final = with_tolerated_stream_scope(params)
|
||||
|
||||
assert tolerated == {
|
||||
"guardrail": "generic_guardrail_api",
|
||||
"mode": "pre_call",
|
||||
"stream_scope": None,
|
||||
}
|
||||
assert params["stream_scope"] == "sometimes"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,68 @@
|
|||
import { Badge } from "@/components/ui/badge";
|
||||
import { GuardrailModeRows } from "./GuardrailModeDisplay";
|
||||
import { GuardrailStreamScopeDetail } from "./StreamScopeFields";
|
||||
import ToolPermissionRulesEditor, { type ToolPermissionConfig } from "./tool_permission/ToolPermissionRulesEditor";
|
||||
|
||||
export const GuardrailReadOnlyDetails = ({
|
||||
guardrailId,
|
||||
guardrailName,
|
||||
displayName,
|
||||
litellmParams,
|
||||
streamScope,
|
||||
defaultOn,
|
||||
piiEntityCount,
|
||||
createdAt,
|
||||
updatedAt,
|
||||
showToolPermission,
|
||||
toolPermissionConfig,
|
||||
}: {
|
||||
guardrailId: string;
|
||||
guardrailName: string;
|
||||
displayName: string;
|
||||
litellmParams: { mode?: unknown; logging_only_scope?: string | null };
|
||||
streamScope: unknown;
|
||||
defaultOn: boolean | undefined;
|
||||
piiEntityCount: number;
|
||||
createdAt: string;
|
||||
updatedAt: string;
|
||||
showToolPermission: boolean;
|
||||
toolPermissionConfig: ToolPermissionConfig;
|
||||
}) => (
|
||||
<div className="space-y-4">
|
||||
<div>
|
||||
<p className="font-medium">Guardrail ID</p>
|
||||
<div className="font-mono">{guardrailId}</div>
|
||||
</div>
|
||||
<div>
|
||||
<p className="font-medium">Guardrail Name</p>
|
||||
<div>{guardrailName || "Unnamed Guardrail"}</div>
|
||||
</div>
|
||||
<div>
|
||||
<p className="font-medium">Provider</p>
|
||||
<div>{displayName}</div>
|
||||
</div>
|
||||
<GuardrailModeRows litellmParams={litellmParams} />
|
||||
<GuardrailStreamScopeDetail raw={streamScope} />
|
||||
<div>
|
||||
<p className="font-medium">Default On</p>
|
||||
<Badge variant={defaultOn ? "secondary" : "outline"}>{defaultOn ? "Yes" : "No"}</Badge>
|
||||
</div>
|
||||
{piiEntityCount > 0 && (
|
||||
<div>
|
||||
<p className="font-medium">PII Protection</p>
|
||||
<div className="mt-2">
|
||||
<Badge variant="secondary">{piiEntityCount} PII entities configured</Badge>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
<div>
|
||||
<p className="font-medium">Created At</p>
|
||||
<div>{createdAt}</div>
|
||||
</div>
|
||||
<div>
|
||||
<p className="font-medium">Last Updated</p>
|
||||
<div>{updatedAt}</div>
|
||||
</div>
|
||||
{showToolPermission && <ToolPermissionRulesEditor value={toolPermissionConfig} disabled />}
|
||||
</div>
|
||||
);
|
||||
|
|
@ -0,0 +1,95 @@
|
|||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { GuardrailField, labelWithHint, type GuardrailFormControl } from "./GuardrailFormField";
|
||||
import {
|
||||
STREAM_SCOPE_OPTIONS,
|
||||
formatGuardrailStreamScope,
|
||||
type GuardrailStreamScope,
|
||||
isGuardrailStreamScope,
|
||||
} from "./guardrail_info_helpers";
|
||||
|
||||
const STREAM_SCOPE_ITEMS = STREAM_SCOPE_OPTIONS.map((option) => ({
|
||||
label: option.label,
|
||||
value: option.value,
|
||||
}));
|
||||
|
||||
const REQUEST_SHAPE_HINT =
|
||||
"Run this guardrail on streaming requests, non-streaming requests, or both, for each selected mode.";
|
||||
|
||||
export const StreamScopeFields = ({
|
||||
modes,
|
||||
value,
|
||||
onChange,
|
||||
}: {
|
||||
modes: string[];
|
||||
value: Record<string, GuardrailStreamScope>;
|
||||
onChange: (next: Record<string, GuardrailStreamScope>) => void;
|
||||
}) => {
|
||||
if (modes.length === 0) return null;
|
||||
|
||||
return (
|
||||
<div className="space-y-3">
|
||||
{modes.map((mode) => {
|
||||
const selected = value[mode] ?? "both";
|
||||
return (
|
||||
<div key={mode} className="space-y-1">
|
||||
<label htmlFor={`stream-scope-${mode}`} className="block text-sm font-medium text-foreground">
|
||||
{mode} applies to
|
||||
</label>
|
||||
<Select
|
||||
items={STREAM_SCOPE_ITEMS}
|
||||
value={selected}
|
||||
onValueChange={(next: GuardrailStreamScope | null) => {
|
||||
if (!isGuardrailStreamScope(next)) return;
|
||||
onChange({ ...value, [mode]: next });
|
||||
}}
|
||||
>
|
||||
<SelectTrigger id={`stream-scope-${mode}`} className="w-full">
|
||||
<SelectValue placeholder="Streaming and non-streaming" />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{STREAM_SCOPE_OPTIONS.map((option) => (
|
||||
<SelectItem key={option.value} value={option.value}>
|
||||
{option.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export const StreamScopeFormField = ({ control, modes }: { control: GuardrailFormControl; modes: string[] }) => (
|
||||
<GuardrailField
|
||||
control={control}
|
||||
name="stream_scope_by_mode"
|
||||
label={labelWithHint("Request shape", REQUEST_SHAPE_HINT)}
|
||||
>
|
||||
{({ value, onChange }) => (
|
||||
<StreamScopeFields
|
||||
modes={modes}
|
||||
value={(value as Record<string, GuardrailStreamScope> | undefined) ?? {}}
|
||||
onChange={onChange}
|
||||
/>
|
||||
)}
|
||||
</GuardrailField>
|
||||
);
|
||||
|
||||
export const GuardrailStreamScopeCaption = ({ raw }: { raw: unknown }) => {
|
||||
const label = formatGuardrailStreamScope(raw);
|
||||
if (!label) return null;
|
||||
return <p className="mt-1 text-sm text-muted-foreground">{label}</p>;
|
||||
};
|
||||
|
||||
export const GuardrailStreamScopeDetail = ({ raw }: { raw: unknown }) => {
|
||||
const label = formatGuardrailStreamScope(raw);
|
||||
if (!label) return null;
|
||||
return (
|
||||
<div>
|
||||
<p className="font-medium">Request shape</p>
|
||||
<div>{label}</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
@ -89,6 +89,23 @@ describe("AddGuardrailForm create payload characterization", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("sends stream_scope when a mode is restricted to streaming requests", async () => {
|
||||
const user = userEvent.setup({ delay: null });
|
||||
renderForm();
|
||||
|
||||
await user.type(await screen.findByLabelText("Guardrail Name"), "my-bedrock");
|
||||
await pickProvider(user, "Bedrock Guardrail");
|
||||
await chooseSelectOption(user, screen.getByLabelText("pre_call applies to"), "Streaming only");
|
||||
await user.type(await screen.findByPlaceholderText("The guardrail id on Bedrock"), "gr-123");
|
||||
await user.click(screen.getByRole("button", { name: "Next" }));
|
||||
await user.click(await screen.findByRole("button", { name: "Create Guardrail" }));
|
||||
|
||||
await waitFor(() => expect(networking.createGuardrailCall).toHaveBeenCalledTimes(1));
|
||||
expect(payload()).toMatchObject({
|
||||
litellm_params: { stream_scope: "streaming" },
|
||||
});
|
||||
});
|
||||
|
||||
it("switches mode from the seeded string to an array once the user touches the multi select", async () => {
|
||||
const user = userEvent.setup({ delay: null });
|
||||
renderForm();
|
||||
|
|
|
|||
|
|
@ -23,11 +23,14 @@ import {
|
|||
shouldRenderContentFilterConfigSettings,
|
||||
shouldRenderLLMJudgeFields,
|
||||
shouldRenderPIIConfigSettings,
|
||||
supportsDirectionalLoggingOnlyScope,
|
||||
streamScopePayload,
|
||||
toModeArray,
|
||||
type GuardrailStreamScope,
|
||||
supportsDirectionalLoggingOnlyScope,
|
||||
type LoggingOnlyScope,
|
||||
type LoggingOnlyScopeChoice,
|
||||
} from "./guardrail_info_helpers";
|
||||
import { StreamScopeFormField } from "./StreamScopeFields";
|
||||
import { Logo } from "@/components/molecules/logo/Logo";
|
||||
import { MultiSelect } from "@/components/shared/MultiSelect";
|
||||
import { FieldGroup } from "@/components/ui/field";
|
||||
|
|
@ -172,6 +175,7 @@ const INITIAL_VALUES: GuardrailFormValues = {
|
|||
logging_only_scope_choice: "default",
|
||||
skip_system_message_choice: "inherit",
|
||||
skip_tool_message_choice: "inherit",
|
||||
stream_scope_by_mode: {},
|
||||
};
|
||||
|
||||
const ALWAYS_ON_ITEMS = [
|
||||
|
|
@ -466,6 +470,14 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
|
|||
guardrail_info: {},
|
||||
};
|
||||
|
||||
const streamScope = streamScopePayload(
|
||||
toModeArray(values.mode),
|
||||
(values.stream_scope_by_mode as Record<string, GuardrailStreamScope> | undefined) ?? {},
|
||||
);
|
||||
if (streamScope !== undefined) {
|
||||
guardrailData.litellm_params.stream_scope = streamScope;
|
||||
}
|
||||
|
||||
const skipForCreate = choiceToSkipSystemForCreate(asSkipChoice(values.skip_system_message_choice));
|
||||
if (skipForCreate !== undefined) {
|
||||
guardrailData.litellm_params.skip_system_message_in_guardrail = skipForCreate;
|
||||
|
|
@ -773,6 +785,8 @@ const AddGuardrailForm: React.FC<AddGuardrailFormProps> = ({ visible, onClose, a
|
|||
)}
|
||||
</GuardrailField>
|
||||
|
||||
<StreamScopeFormField control={form.control} modes={toModeArray(form.watch("mode"))} />
|
||||
|
||||
<GuardrailField
|
||||
control={form.control}
|
||||
name="default_on"
|
||||
|
|
|
|||
|
|
@ -68,6 +68,62 @@ describe("CustomCodeModal", () => {
|
|||
expect(screen.getByRole("button", { name: /update guardrail/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should open without crashing for a tag-scoped mode and show it read-only", async () => {
|
||||
renderModal({
|
||||
editData: {
|
||||
guardrail_id: "g-tag",
|
||||
guardrail_name: "tag-mode-guardrail",
|
||||
litellm_params: {
|
||||
mode: { tags: { "team-a": "pre_call" }, default: "post_call" },
|
||||
default_on: true,
|
||||
custom_code: "def apply_guardrail(): pass",
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
expect(await screen.findByText("Edit Custom Guardrail")).toBeInTheDocument();
|
||||
const modeInput = screen.getByLabelText("Mode (tag-scoped)");
|
||||
expect(modeInput).toBeDisabled();
|
||||
expect(modeInput).toHaveValue("post_call, pre_call (tag-based)");
|
||||
expect(screen.getByText("Mode (tag-scoped, read-only)")).toBeInTheDocument();
|
||||
expect(screen.queryByText(/applies to/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should omit mode and stream_scope from the update payload for a tag-scoped guardrail", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderModal({
|
||||
editData: {
|
||||
guardrail_id: "g-tag",
|
||||
guardrail_name: "tag-mode-guardrail",
|
||||
litellm_params: {
|
||||
mode: { tags: { "team-a": "pre_call" }, default: "post_call" },
|
||||
default_on: true,
|
||||
custom_code: "def apply_guardrail(): pass",
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const nameInput = await screen.findByDisplayValue("tag-mode-guardrail");
|
||||
await user.clear(nameInput);
|
||||
await user.type(nameInput, "renamed-guardrail");
|
||||
await user.click(screen.getByRole("button", { name: /update guardrail/i }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockUpdate).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
const [token, guardrailId, payload] = mockUpdate.mock.calls[0] as [
|
||||
string,
|
||||
string,
|
||||
Record<string, Record<string, unknown>>,
|
||||
];
|
||||
expect(token).toBe("test-token");
|
||||
expect(guardrailId).toBe("g-tag");
|
||||
expect(payload.guardrail_name).toBe("renamed-guardrail");
|
||||
expect(payload.litellm_params.custom_code).toBe("def apply_guardrail(): pass");
|
||||
expect(payload.litellm_params).not.toHaveProperty("mode");
|
||||
expect(payload.litellm_params).not.toHaveProperty("stream_scope");
|
||||
});
|
||||
|
||||
it("should keep save disabled until a guardrail name is entered", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderModal();
|
||||
|
|
@ -111,8 +167,7 @@ describe("CustomCodeModal", () => {
|
|||
|
||||
expect(await screen.findByDisplayValue(/async def apply_guardrail/)).toBeInTheDocument();
|
||||
|
||||
const comboboxes = screen.getAllByRole("combobox");
|
||||
await user.click(comboboxes[comboboxes.length - 1]);
|
||||
await user.click(screen.getByRole("combobox", { name: "Template" }));
|
||||
const options = await screen.findAllByText("Block SSN");
|
||||
await user.click(options[options.length - 1]);
|
||||
|
||||
|
|
|
|||
|
|
@ -37,152 +37,30 @@ import {
|
|||
import { Switch } from "@/components/ui/switch";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
|
||||
|
||||
// Code templates
|
||||
const CODE_TEMPLATES = {
|
||||
empty: {
|
||||
name: "Empty Template",
|
||||
code: `async def apply_guardrail(inputs, request_data, input_type):
|
||||
# inputs: {texts, images, tools, tool_calls, structured_messages, model}
|
||||
# request_data: {model, user_id, team_id, end_user_id, metadata}
|
||||
# input_type: "request" or "response"
|
||||
return allow()`,
|
||||
},
|
||||
blockSSN: {
|
||||
name: "Block SSN",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
for text in inputs["texts"]:
|
||||
if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"):
|
||||
return block("SSN detected")
|
||||
return allow()`,
|
||||
},
|
||||
redactEmail: {
|
||||
name: "Redact Emails",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}"
|
||||
modified = []
|
||||
for text in inputs["texts"]:
|
||||
modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]"))
|
||||
return modify(texts=modified)`,
|
||||
},
|
||||
blockSQL: {
|
||||
name: "Block SQL Injection",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
if input_type != "request":
|
||||
return allow()
|
||||
for text in inputs["texts"]:
|
||||
if contains_code_language(text, ["sql"]):
|
||||
return block("SQL code not allowed")
|
||||
return allow()`,
|
||||
},
|
||||
validateJSON: {
|
||||
name: "Validate JSON",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
if input_type != "response":
|
||||
return allow()
|
||||
|
||||
schema = {"type": "object", "required": ["name", "value"]}
|
||||
|
||||
for text in inputs["texts"]:
|
||||
obj = json_parse(text)
|
||||
if obj is None:
|
||||
return block("Invalid JSON response")
|
||||
if not json_schema_valid(obj, schema):
|
||||
return block("Response missing required fields")
|
||||
return allow()`,
|
||||
},
|
||||
externalAPI: {
|
||||
name: "External API Check (async)",
|
||||
code: `async def apply_guardrail(inputs, request_data, input_type):
|
||||
# Call an external moderation API (async for non-blocking)
|
||||
for text in inputs["texts"]:
|
||||
response = await http_post(
|
||||
"https://api.example.com/moderate",
|
||||
body={"text": text, "user_id": request_data["user_id"]},
|
||||
headers={"Authorization": "Bearer YOUR_API_KEY"},
|
||||
timeout=10
|
||||
)
|
||||
|
||||
if not response["success"]:
|
||||
# API call failed, allow by default or block
|
||||
return allow()
|
||||
|
||||
if response["body"].get("flagged"):
|
||||
return block(response["body"].get("reason", "Content flagged"))
|
||||
|
||||
return allow()`,
|
||||
},
|
||||
};
|
||||
|
||||
// Available primitives organized by category
|
||||
const PRIMITIVES = {
|
||||
"Return Values": [
|
||||
{ name: "allow()", desc: "Let request/response through" },
|
||||
{ name: "block(reason)", desc: "Reject with message" },
|
||||
{ name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" },
|
||||
{ name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" },
|
||||
],
|
||||
"HTTP Requests (async)": [
|
||||
{ name: "await http_request(url, method, headers, body)", desc: "Make async HTTP request" },
|
||||
{ name: "await http_get(url, headers)", desc: "Async GET request" },
|
||||
{ name: "await http_post(url, body, headers)", desc: "Async POST request" },
|
||||
],
|
||||
"Regex Functions": [
|
||||
{ name: "regex_match(text, pattern)", desc: "Returns True if pattern found" },
|
||||
{ name: "regex_replace(text, pattern, replacement)", desc: "Replace all matches" },
|
||||
{ name: "regex_find_all(text, pattern)", desc: "Return list of matches" },
|
||||
],
|
||||
"JSON Functions": [
|
||||
{ name: "json_parse(text)", desc: "Parse JSON string, returns None on error" },
|
||||
{ name: "json_stringify(obj)", desc: "Convert to JSON string" },
|
||||
{ name: "json_schema_valid(obj, schema)", desc: "Validate against JSON schema" },
|
||||
],
|
||||
"URL Functions": [
|
||||
{ name: "extract_urls(text)", desc: "Extract all URLs from text" },
|
||||
{ name: "is_valid_url(url)", desc: "Check if URL is valid" },
|
||||
{ name: "all_urls_valid(text)", desc: "Check all URLs in text are valid" },
|
||||
],
|
||||
"Code Detection": [
|
||||
{ name: "detect_code(text)", desc: "Returns True if code detected" },
|
||||
{ name: "detect_code_languages(text)", desc: "Returns list of detected languages" },
|
||||
{ name: 'contains_code_language(text, ["sql"])', desc: "Check for specific languages" },
|
||||
],
|
||||
"Text Utilities": [
|
||||
{ name: "contains(text, substring)", desc: "Check if substring exists" },
|
||||
{ name: "contains_any(text, [substr1, substr2])", desc: "Check if any substring exists" },
|
||||
{ name: "word_count(text)", desc: "Count words" },
|
||||
{ name: "char_count(text)", desc: "Count characters" },
|
||||
{ name: "lower(text) / upper(text) / trim(text)", desc: "String transforms" },
|
||||
],
|
||||
};
|
||||
|
||||
const MODE_OPTIONS = [
|
||||
{ value: "pre_call", label: "pre_call (Request)" },
|
||||
{ value: "post_call", label: "post_call (Response)" },
|
||||
{ value: "during_call", label: "during_call (Parallel)" },
|
||||
{ value: "logging_only", label: "logging_only" },
|
||||
{ value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" },
|
||||
{ value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" },
|
||||
{ value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" },
|
||||
];
|
||||
|
||||
const TEMPLATE_ITEMS = Object.entries(CODE_TEMPLATES).map(([key, template]) => ({
|
||||
value: key,
|
||||
label: template.name,
|
||||
}));
|
||||
|
||||
type ModeOption = (typeof MODE_OPTIONS)[number];
|
||||
|
||||
const MODE_OPTION_BY_VALUE: Record<string, ModeOption> = Object.fromEntries(
|
||||
MODE_OPTIONS.map((option) => [option.value, option]),
|
||||
);
|
||||
import { StreamScopeFields } from "../StreamScopeFields";
|
||||
import {
|
||||
formatGuardrailMode,
|
||||
streamScopeByModeFromConfig,
|
||||
streamScopeForUpdate,
|
||||
streamScopePayload,
|
||||
type GuardrailStreamScope,
|
||||
} from "../guardrail_info_helpers";
|
||||
import {
|
||||
CODE_TEMPLATES,
|
||||
MODE_OPTION_BY_VALUE,
|
||||
MODE_OPTIONS,
|
||||
PRIMITIVES,
|
||||
TEMPLATE_ITEMS,
|
||||
type ModeOption,
|
||||
} from "./custom_code_catalog";
|
||||
|
||||
// Data for editing an existing guardrail
|
||||
|
||||
export interface EditGuardrailData {
|
||||
guardrail_id: string;
|
||||
guardrail_name: string;
|
||||
litellm_params: {
|
||||
mode?: string | string[];
|
||||
mode?: string | string[] | Record<string, unknown>;
|
||||
default_on?: boolean;
|
||||
custom_code?: string;
|
||||
logging_only_scope?: LoggingOnlyScope | null;
|
||||
|
|
@ -204,6 +82,7 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
const isEditMode = !!editData;
|
||||
const [guardrailName, setGuardrailName] = useState("");
|
||||
const [mode, setMode] = useState<string[]>(["pre_call"]);
|
||||
const [streamScopeByMode, setStreamScopeByMode] = useState<Record<string, GuardrailStreamScope>>({});
|
||||
const [loggingOnlyScopeChoice, setLoggingOnlyScopeChoice] = useState<LoggingOnlyScopeChoice>("default");
|
||||
const [defaultOn, setDefaultOn] = useState(false);
|
||||
const [selectedTemplate, setSelectedTemplate] = useState<string>("empty");
|
||||
|
|
@ -315,11 +194,14 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
setCode(CODE_TEMPLATES[templateKey as keyof typeof CODE_TEMPLATES].code);
|
||||
};
|
||||
|
||||
// Normalize mode from API (string or string[]) to string[]
|
||||
const normalizeMode = (m: string | string[] | undefined): string[] => {
|
||||
// Normalize mode from API (string or string[]) to string[].
|
||||
// A tag-scoped mode dict ({ tags, default }) is managed outside this editor, so it
|
||||
// contributes no editable modes and is displayed read-only instead.
|
||||
const normalizeMode = (m: string | string[] | Record<string, unknown> | undefined): string[] => {
|
||||
if (m === undefined || m === null) return ["pre_call"];
|
||||
if (Array.isArray(m)) return m.length ? m : ["pre_call"];
|
||||
return [m];
|
||||
if (typeof m === "string") return [m];
|
||||
return [];
|
||||
};
|
||||
|
||||
// Reset form when modal opens or editData changes
|
||||
|
|
@ -329,6 +211,12 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
// Edit mode: populate with existing data
|
||||
setGuardrailName(editData.guardrail_name || "");
|
||||
setMode(normalizeMode(editData.litellm_params?.mode));
|
||||
setStreamScopeByMode(
|
||||
streamScopeByModeFromConfig(
|
||||
editData.litellm_params?.stream_scope,
|
||||
normalizeMode(editData.litellm_params?.mode),
|
||||
),
|
||||
);
|
||||
setLoggingOnlyScopeChoice(loggingOnlyScopeToChoice(editData.litellm_params?.logging_only_scope));
|
||||
setDefaultOn(editData.litellm_params?.default_on || false);
|
||||
setCode(editData.litellm_params?.custom_code || CODE_TEMPLATES.empty.code);
|
||||
|
|
@ -337,6 +225,7 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
// Create mode: reset to defaults
|
||||
setGuardrailName("");
|
||||
setMode(["pre_call"]);
|
||||
setStreamScopeByMode({});
|
||||
setLoggingOnlyScopeChoice("default");
|
||||
setDefaultOn(false);
|
||||
setSelectedTemplate("empty");
|
||||
|
|
@ -411,11 +300,21 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
if (defaultOn !== editData.litellm_params?.default_on) {
|
||||
updateData.litellm_params.default_on = defaultOn;
|
||||
}
|
||||
const nextStreamScope = streamScopeForUpdate(
|
||||
mode,
|
||||
streamScopeByMode,
|
||||
editData.litellm_params?.stream_scope,
|
||||
existingMode,
|
||||
);
|
||||
if (nextStreamScope !== undefined) {
|
||||
updateData.litellm_params.stream_scope = nextStreamScope;
|
||||
}
|
||||
|
||||
await updateGuardrailCall(accessToken, editData.guardrail_id, updateData);
|
||||
toast.success("Custom code guardrail updated successfully");
|
||||
} else {
|
||||
// Create new guardrail
|
||||
const streamScope = streamScopePayload(mode, streamScopeByMode);
|
||||
const guardrailData = {
|
||||
guardrail_name: guardrailName,
|
||||
litellm_params: {
|
||||
|
|
@ -423,6 +322,7 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
mode: mode,
|
||||
default_on: defaultOn,
|
||||
custom_code: code,
|
||||
...(streamScope !== undefined ? { stream_scope: streamScope } : {}),
|
||||
...getCustomCodeLoggingOnlyScopeCreate(mode, loggingOnlyScopeChoice),
|
||||
},
|
||||
guardrail_info: {},
|
||||
|
|
@ -511,6 +411,11 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
|
||||
const lineCount = code.split("\n").length;
|
||||
const selectedModeOptions = mode.map((value) => MODE_OPTION_BY_VALUE[value]).filter(Boolean);
|
||||
const rawEditMode = editData?.litellm_params?.mode;
|
||||
const tagScopedModeLabel =
|
||||
rawEditMode !== null && typeof rawEditMode === "object" && !Array.isArray(rawEditMode)
|
||||
? formatGuardrailMode(rawEditMode) || "-"
|
||||
: null;
|
||||
|
||||
return (
|
||||
<Dialog open={visible} onOpenChange={(open) => !open && onClose()}>
|
||||
|
|
@ -533,32 +438,38 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
/>
|
||||
</div>
|
||||
<div className="w-[280px]">
|
||||
<label className="mb-1 block text-xs font-medium text-muted-foreground">Mode (can select multiple)</label>
|
||||
<Combobox
|
||||
items={MODE_OPTIONS}
|
||||
value={selectedModeOptions}
|
||||
onValueChange={(options: ModeOption[]) => setMode(options.map((option) => option.value))}
|
||||
multiple
|
||||
>
|
||||
<ComboboxChips render={<div ref={anchor} />} className="w-full">
|
||||
{selectedModeOptions.map((option) => (
|
||||
<ComboboxChip key={option.value} aria-label={option.label}>
|
||||
{option.label}
|
||||
</ComboboxChip>
|
||||
))}
|
||||
<ComboboxChipsInput placeholder={mode.length === 0 ? "Select modes" : undefined} />
|
||||
</ComboboxChips>
|
||||
<ComboboxContent anchor={anchor}>
|
||||
<ComboboxEmpty>No matching modes</ComboboxEmpty>
|
||||
<ComboboxList>
|
||||
{(option: ModeOption) => (
|
||||
<ComboboxItem key={option.value} value={option}>
|
||||
<label className="mb-1 block text-xs font-medium text-muted-foreground">
|
||||
{tagScopedModeLabel ? "Mode (tag-scoped, read-only)" : "Mode (can select multiple)"}
|
||||
</label>
|
||||
{tagScopedModeLabel ? (
|
||||
<Input value={tagScopedModeLabel} disabled aria-label="Mode (tag-scoped)" />
|
||||
) : (
|
||||
<Combobox
|
||||
items={MODE_OPTIONS}
|
||||
value={selectedModeOptions}
|
||||
onValueChange={(options: ModeOption[]) => setMode(options.map((option) => option.value))}
|
||||
multiple
|
||||
>
|
||||
<ComboboxChips render={<div ref={anchor} />} className="w-full">
|
||||
{selectedModeOptions.map((option) => (
|
||||
<ComboboxChip key={option.value} aria-label={option.label}>
|
||||
{option.label}
|
||||
</ComboboxItem>
|
||||
)}
|
||||
</ComboboxList>
|
||||
</ComboboxContent>
|
||||
</Combobox>
|
||||
</ComboboxChip>
|
||||
))}
|
||||
<ComboboxChipsInput placeholder={mode.length === 0 ? "Select modes" : undefined} />
|
||||
</ComboboxChips>
|
||||
<ComboboxContent anchor={anchor}>
|
||||
<ComboboxEmpty>No matching modes</ComboboxEmpty>
|
||||
<ComboboxList>
|
||||
{(option: ModeOption) => (
|
||||
<ComboboxItem key={option.value} value={option}>
|
||||
{option.label}
|
||||
</ComboboxItem>
|
||||
)}
|
||||
</ComboboxList>
|
||||
</ComboboxContent>
|
||||
</Combobox>
|
||||
)}
|
||||
</div>
|
||||
{mode.includes("logging_only") && (
|
||||
<CustomCodeLoggingOnlyScopeSelect value={loggingOnlyScopeChoice} onChange={setLoggingOnlyScopeChoice} />
|
||||
|
|
@ -600,6 +511,11 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({ visible, onClose, onS
|
|||
<Switch checked={defaultOn} onCheckedChange={setDefaultOn} aria-label="Default On" />
|
||||
</div>
|
||||
</div>
|
||||
{mode.length > 0 && (
|
||||
<div className="border-b border-border py-4">
|
||||
<StreamScopeFields modes={mode} value={streamScopeByMode} onChange={setStreamScopeByMode} />
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Main Content */}
|
||||
<div className="mt-4 flex gap-6">
|
||||
|
|
|
|||
|
|
@ -0,0 +1,136 @@
|
|||
export const CODE_TEMPLATES = {
|
||||
empty: {
|
||||
name: "Empty Template",
|
||||
code: `async def apply_guardrail(inputs, request_data, input_type):
|
||||
# inputs: {texts, images, tools, tool_calls, structured_messages, model}
|
||||
# request_data: {model, user_id, team_id, end_user_id, metadata}
|
||||
# input_type: "request" or "response"
|
||||
return allow()`,
|
||||
},
|
||||
blockSSN: {
|
||||
name: "Block SSN",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
for text in inputs["texts"]:
|
||||
if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"):
|
||||
return block("SSN detected")
|
||||
return allow()`,
|
||||
},
|
||||
redactEmail: {
|
||||
name: "Redact Emails",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}"
|
||||
modified = []
|
||||
for text in inputs["texts"]:
|
||||
modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]"))
|
||||
return modify(texts=modified)`,
|
||||
},
|
||||
blockSQL: {
|
||||
name: "Block SQL Injection",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
if input_type != "request":
|
||||
return allow()
|
||||
for text in inputs["texts"]:
|
||||
if contains_code_language(text, ["sql"]):
|
||||
return block("SQL code not allowed")
|
||||
return allow()`,
|
||||
},
|
||||
validateJSON: {
|
||||
name: "Validate JSON",
|
||||
code: `def apply_guardrail(inputs, request_data, input_type):
|
||||
if input_type != "response":
|
||||
return allow()
|
||||
|
||||
schema = {"type": "object", "required": ["name", "value"]}
|
||||
|
||||
for text in inputs["texts"]:
|
||||
obj = json_parse(text)
|
||||
if obj is None:
|
||||
return block("Invalid JSON response")
|
||||
if not json_schema_valid(obj, schema):
|
||||
return block("Response missing required fields")
|
||||
return allow()`,
|
||||
},
|
||||
externalAPI: {
|
||||
name: "External API Check (async)",
|
||||
code: `async def apply_guardrail(inputs, request_data, input_type):
|
||||
# Call an external moderation API (async for non-blocking)
|
||||
for text in inputs["texts"]:
|
||||
response = await http_post(
|
||||
"https://api.example.com/moderate",
|
||||
body={"text": text, "user_id": request_data["user_id"]},
|
||||
headers={"Authorization": "Bearer YOUR_API_KEY"},
|
||||
timeout=10
|
||||
)
|
||||
|
||||
if not response["success"]:
|
||||
# API call failed, allow by default or block
|
||||
return allow()
|
||||
|
||||
if response["body"].get("flagged"):
|
||||
return block(response["body"].get("reason", "Content flagged"))
|
||||
|
||||
return allow()`,
|
||||
},
|
||||
};
|
||||
|
||||
export const PRIMITIVES = {
|
||||
"Return Values": [
|
||||
{ name: "allow()", desc: "Let request/response through" },
|
||||
{ name: "block(reason)", desc: "Reject with message" },
|
||||
{ name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" },
|
||||
{ name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" },
|
||||
],
|
||||
"HTTP Requests (async)": [
|
||||
{ name: "await http_request(url, method, headers, body)", desc: "Make async HTTP request" },
|
||||
{ name: "await http_get(url, headers)", desc: "Async GET request" },
|
||||
{ name: "await http_post(url, body, headers)", desc: "Async POST request" },
|
||||
],
|
||||
"Regex Functions": [
|
||||
{ name: "regex_match(text, pattern)", desc: "Returns True if pattern found" },
|
||||
{ name: "regex_replace(text, pattern, replacement)", desc: "Replace all matches" },
|
||||
{ name: "regex_find_all(text, pattern)", desc: "Return list of matches" },
|
||||
],
|
||||
"JSON Functions": [
|
||||
{ name: "json_parse(text)", desc: "Parse JSON string, returns None on error" },
|
||||
{ name: "json_stringify(obj)", desc: "Convert to JSON string" },
|
||||
{ name: "json_schema_valid(obj, schema)", desc: "Validate against JSON schema" },
|
||||
],
|
||||
"URL Functions": [
|
||||
{ name: "extract_urls(text)", desc: "Extract all URLs from text" },
|
||||
{ name: "is_valid_url(url)", desc: "Check if URL is valid" },
|
||||
{ name: "all_urls_valid(text)", desc: "Check all URLs in text are valid" },
|
||||
],
|
||||
"Code Detection": [
|
||||
{ name: "detect_code(text)", desc: "Returns True if code detected" },
|
||||
{ name: "detect_code_languages(text)", desc: "Returns list of detected languages" },
|
||||
{ name: 'contains_code_language(text, ["sql"])', desc: "Check for specific languages" },
|
||||
],
|
||||
"Text Utilities": [
|
||||
{ name: "contains(text, substring)", desc: "Check if substring exists" },
|
||||
{ name: "contains_any(text, [substr1, substr2])", desc: "Check if any substring exists" },
|
||||
{ name: "word_count(text)", desc: "Count words" },
|
||||
{ name: "char_count(text)", desc: "Count characters" },
|
||||
{ name: "lower(text) / upper(text) / trim(text)", desc: "String transforms" },
|
||||
],
|
||||
};
|
||||
|
||||
export const MODE_OPTIONS = [
|
||||
{ value: "pre_call", label: "pre_call (Request)" },
|
||||
{ value: "post_call", label: "post_call (Response)" },
|
||||
{ value: "during_call", label: "during_call (Parallel)" },
|
||||
{ value: "logging_only", label: "logging_only" },
|
||||
{ value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" },
|
||||
{ value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" },
|
||||
{ value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" },
|
||||
];
|
||||
|
||||
export const TEMPLATE_ITEMS = Object.entries(CODE_TEMPLATES).map(([key, template]) => ({
|
||||
value: key,
|
||||
label: template.name,
|
||||
}));
|
||||
|
||||
export type ModeOption = (typeof MODE_OPTIONS)[number];
|
||||
|
||||
export const MODE_OPTION_BY_VALUE: Record<string, ModeOption> = Object.fromEntries(
|
||||
MODE_OPTIONS.map((option) => [option.value, option]),
|
||||
);
|
||||
|
|
@ -33,7 +33,8 @@ import {
|
|||
} from "./GuardrailFormField";
|
||||
import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager";
|
||||
import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal";
|
||||
import { GuardrailModeCard, GuardrailModeRows } from "./GuardrailModeDisplay";
|
||||
import { GuardrailModeCard } from "./GuardrailModeDisplay";
|
||||
import { GuardrailReadOnlyDetails } from "./GuardrailReadOnlyDetails";
|
||||
import {
|
||||
getLoggingOnlyScopeUpdate,
|
||||
getGuardrailLogoAndName,
|
||||
|
|
@ -41,10 +42,15 @@ import {
|
|||
loggingOnlyScopeToChoice,
|
||||
skipSystemMessageToChoice,
|
||||
skipToolMessageToChoice,
|
||||
streamScopeByModeFromConfig,
|
||||
streamScopeForUpdate,
|
||||
supportsDirectionalLoggingOnlyScope,
|
||||
toModeArray,
|
||||
type SkipSystemMessageChoice,
|
||||
type SkipToolMessageChoice,
|
||||
type GuardrailStreamScope,
|
||||
} from "./guardrail_info_helpers";
|
||||
import { GuardrailStreamScopeCaption, StreamScopeFormField } from "./StreamScopeFields";
|
||||
import GuardrailOptionalParams from "./guardrail_optional_params";
|
||||
import GuardrailProviderFields from "./guardrail_provider_fields";
|
||||
import PiiConfiguration from "./pii_configuration";
|
||||
|
|
@ -236,6 +242,13 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
|
|||
"skip_tool_message_choice",
|
||||
skipToolMessageToChoice(guardrailData.litellm_params?.skip_tool_message_in_guardrail),
|
||||
);
|
||||
form.setValue(
|
||||
"stream_scope_by_mode",
|
||||
streamScopeByModeFromConfig(
|
||||
guardrailData.litellm_params?.stream_scope,
|
||||
toModeArray(guardrailData.litellm_params?.mode),
|
||||
),
|
||||
);
|
||||
form.setValue(
|
||||
"guardrail_info",
|
||||
guardrailData.guardrail_info ? JSON.stringify(guardrailData.guardrail_info, null, 2) : "",
|
||||
|
|
@ -304,6 +317,16 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
|
|||
updateData.litellm_params.default_on = values.default_on;
|
||||
}
|
||||
|
||||
const modes = toModeArray(guardrailData.litellm_params?.mode);
|
||||
const nextStreamScope = streamScopeForUpdate(
|
||||
modes,
|
||||
(values.stream_scope_by_mode as Record<string, GuardrailStreamScope> | undefined) ?? {},
|
||||
guardrailData.litellm_params?.stream_scope,
|
||||
);
|
||||
if (nextStreamScope !== undefined) {
|
||||
updateData.litellm_params.stream_scope = nextStreamScope;
|
||||
}
|
||||
|
||||
const prevSkipChoice = skipSystemMessageToChoice(guardrailData.litellm_params?.skip_system_message_in_guardrail);
|
||||
const nextSkipChoice = values.skip_system_message_choice as SkipSystemMessageChoice | undefined;
|
||||
if (nextSkipChoice !== undefined && nextSkipChoice !== prevSkipChoice) {
|
||||
|
|
@ -566,6 +589,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
|
|||
</Card>
|
||||
|
||||
<GuardrailModeCard litellmParams={guardrailData.litellm_params} />
|
||||
<GuardrailStreamScopeCaption raw={guardrailData.litellm_params?.stream_scope} />
|
||||
|
||||
<Card className="block p-6">
|
||||
<p>Created At</p>
|
||||
|
|
@ -723,6 +747,11 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
|
|||
)}
|
||||
</GuardrailField>
|
||||
|
||||
<StreamScopeFormField
|
||||
control={form.control}
|
||||
modes={toModeArray(guardrailData.litellm_params?.mode)}
|
||||
/>
|
||||
|
||||
<GuardrailField
|
||||
control={form.control}
|
||||
name="skip_system_message_choice"
|
||||
|
|
@ -847,53 +876,19 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
|
|||
</form>
|
||||
</TooltipProvider>
|
||||
) : (
|
||||
<div className="space-y-4">
|
||||
<div>
|
||||
<p className="font-medium">Guardrail ID</p>
|
||||
<div className="font-mono">{guardrailData.guardrail_id}</div>
|
||||
</div>
|
||||
<div>
|
||||
<p className="font-medium">Guardrail Name</p>
|
||||
<div>{guardrailData.guardrail_name || "Unnamed Guardrail"}</div>
|
||||
</div>
|
||||
<div>
|
||||
<p className="font-medium">Provider</p>
|
||||
<div>{displayName}</div>
|
||||
</div>
|
||||
<GuardrailModeRows litellmParams={guardrailData.litellm_params} />
|
||||
<div>
|
||||
<p className="font-medium">Default On</p>
|
||||
<Badge variant={guardrailData.litellm_params?.default_on ? "secondary" : "outline"}>
|
||||
{guardrailData.litellm_params?.default_on ? "Yes" : "No"}
|
||||
</Badge>
|
||||
</div>
|
||||
|
||||
{guardrailData.litellm_params?.pii_entities_config &&
|
||||
Object.keys(guardrailData.litellm_params.pii_entities_config).length > 0 && (
|
||||
<div>
|
||||
<p className="font-medium">PII Protection</p>
|
||||
<div className="mt-2">
|
||||
<Badge variant="secondary">
|
||||
{Object.keys(guardrailData.litellm_params.pii_entities_config).length} PII entities
|
||||
configured
|
||||
</Badge>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div>
|
||||
<p className="font-medium">Created At</p>
|
||||
<div>{formatDate(guardrailData.created_at)}</div>
|
||||
</div>
|
||||
<div>
|
||||
<p className="font-medium">Last Updated</p>
|
||||
<div>{formatDate(guardrailData.updated_at)}</div>
|
||||
</div>
|
||||
|
||||
{guardrailData.litellm_params?.guardrail === "tool_permission" && (
|
||||
<ToolPermissionRulesEditor value={toolPermissionConfig} disabled />
|
||||
)}
|
||||
</div>
|
||||
<GuardrailReadOnlyDetails
|
||||
guardrailId={guardrailData.guardrail_id}
|
||||
guardrailName={guardrailData.guardrail_name}
|
||||
displayName={displayName}
|
||||
litellmParams={guardrailData.litellm_params}
|
||||
streamScope={guardrailData.litellm_params?.stream_scope}
|
||||
defaultOn={guardrailData.litellm_params?.default_on}
|
||||
piiEntityCount={Object.keys(guardrailData.litellm_params?.pii_entities_config || {}).length}
|
||||
createdAt={formatDate(guardrailData.created_at)}
|
||||
updatedAt={formatDate(guardrailData.updated_at)}
|
||||
showToolPermission={guardrailData.litellm_params?.guardrail === "tool_permission"}
|
||||
toolPermissionConfig={toolPermissionConfig}
|
||||
/>
|
||||
)}
|
||||
</Card>
|
||||
</TabsContent>
|
||||
|
|
|
|||
|
|
@ -15,6 +15,11 @@ import {
|
|||
skipToolMessageToChoice,
|
||||
choiceToSkipToolForCreate,
|
||||
formatGuardrailMode,
|
||||
formatGuardrailStreamScope,
|
||||
streamScopeByModeFromConfig,
|
||||
streamScopeForMode,
|
||||
streamScopeForUpdate,
|
||||
streamScopePayload,
|
||||
loggingOnlyScopeToChoice,
|
||||
choiceToLoggingOnlyScope,
|
||||
getLoggingOnlyScopeUpdate,
|
||||
|
|
@ -369,4 +374,56 @@ describe("guardrail_info_helpers", () => {
|
|||
expect(choiceToSkipToolForCreate("no")).toBe(false);
|
||||
});
|
||||
});
|
||||
|
||||
describe("stream_scope helpers", () => {
|
||||
it("treats omitted config as both for every mode", () => {
|
||||
expect(streamScopeForMode(undefined, "pre_call")).toBe("both");
|
||||
expect(streamScopeByModeFromConfig(undefined, ["pre_call", "post_call"])).toEqual({
|
||||
pre_call: "both",
|
||||
post_call: "both",
|
||||
});
|
||||
});
|
||||
|
||||
it("applies a scalar to every selected mode and omits both-only payloads", () => {
|
||||
expect(streamScopeForMode("streaming", "post_call")).toBe("streaming");
|
||||
expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "both", post_call: "both" })).toBeUndefined();
|
||||
expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "streaming", post_call: "streaming" })).toBe(
|
||||
"streaming",
|
||||
);
|
||||
});
|
||||
|
||||
it("keeps a mixed map instead of collapsing it to a scalar", () => {
|
||||
expect(streamScopePayload(["pre_call", "post_call"], { pre_call: "both", post_call: "streaming" })).toEqual({
|
||||
post_call: "streaming",
|
||||
});
|
||||
});
|
||||
|
||||
it("formats scalar and per-mode stream scopes for display", () => {
|
||||
expect(formatGuardrailStreamScope(undefined)).toBe("");
|
||||
expect(formatGuardrailStreamScope("both")).toBe("Streaming and non-streaming");
|
||||
expect(formatGuardrailStreamScope("non_streaming")).toBe("Non-streaming only");
|
||||
expect(formatGuardrailStreamScope({ post_call: "streaming", pre_call: "both" })).toBe(
|
||||
"post_call: Streaming only, pre_call: Streaming and non-streaming",
|
||||
);
|
||||
});
|
||||
|
||||
it("emits both on update only when a prior restriction is cleared", () => {
|
||||
expect(streamScopeForUpdate(["post_call"], { post_call: "streaming" }, undefined)).toBe("streaming");
|
||||
expect(streamScopeForUpdate(["post_call"], { post_call: "both" }, "streaming")).toBe("both");
|
||||
expect(streamScopeForUpdate(["post_call"], { post_call: "both" }, undefined)).toBeUndefined();
|
||||
});
|
||||
|
||||
it("keeps stored restrictions for modes outside the current selection", () => {
|
||||
expect(
|
||||
streamScopeForUpdate(
|
||||
["pre_call"],
|
||||
{ pre_call: "streaming" },
|
||||
{ pre_call: "streaming", post_call: "non_streaming" },
|
||||
),
|
||||
).toBeUndefined();
|
||||
expect(
|
||||
streamScopeForUpdate(["pre_call"], { pre_call: "both" }, { pre_call: "streaming", post_call: "non_streaming" }),
|
||||
).toEqual({ post_call: "non_streaming" });
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -343,3 +343,86 @@ export function choiceToSkipToolForCreate(choice: SkipToolMessageChoice | undefi
|
|||
if (choice === "no") return false;
|
||||
return undefined;
|
||||
}
|
||||
|
||||
export const GUARDRAIL_STREAM_SCOPES = ["both", "streaming", "non_streaming"] as const;
|
||||
export type GuardrailStreamScope = (typeof GUARDRAIL_STREAM_SCOPES)[number];
|
||||
|
||||
export const STREAM_SCOPE_OPTIONS: { value: GuardrailStreamScope; label: string }[] = [
|
||||
{ value: "both", label: "Streaming and non-streaming" },
|
||||
{ value: "streaming", label: "Streaming only" },
|
||||
{ value: "non_streaming", label: "Non-streaming only" },
|
||||
];
|
||||
|
||||
export const isGuardrailStreamScope = (value: unknown): value is GuardrailStreamScope =>
|
||||
value === "both" || value === "streaming" || value === "non_streaming";
|
||||
|
||||
export const streamScopeForMode = (raw: unknown, mode: string): GuardrailStreamScope => {
|
||||
if (isGuardrailStreamScope(raw)) return raw;
|
||||
if (raw !== null && typeof raw === "object" && !Array.isArray(raw)) {
|
||||
const value = (raw as Record<string, unknown>)[mode];
|
||||
if (isGuardrailStreamScope(value)) return value;
|
||||
}
|
||||
return "both";
|
||||
};
|
||||
|
||||
export const streamScopeByModeFromConfig = (raw: unknown, modes: string[]): Record<string, GuardrailStreamScope> =>
|
||||
Object.fromEntries(modes.map((mode) => [mode, streamScopeForMode(raw, mode)]));
|
||||
|
||||
export const streamScopePayload = (
|
||||
modes: string[],
|
||||
scopes: Record<string, GuardrailStreamScope>,
|
||||
): GuardrailStreamScope | Record<string, GuardrailStreamScope> | undefined => {
|
||||
const perMode: Record<string, GuardrailStreamScope> = Object.fromEntries(
|
||||
modes.map((mode) => [mode, scopes[mode] ?? "both"]),
|
||||
);
|
||||
const values = Object.values(perMode);
|
||||
if (values.length === 0 || values.every((scope) => scope === "both")) return undefined;
|
||||
const unique = new Set(values);
|
||||
if (unique.size === 1) return values[0];
|
||||
return Object.fromEntries(Object.entries(perMode).filter((entry) => entry[1] !== "both"));
|
||||
};
|
||||
|
||||
export const formatGuardrailStreamScope = (raw: unknown): string => {
|
||||
if (isGuardrailStreamScope(raw)) {
|
||||
return STREAM_SCOPE_OPTIONS.find((option) => option.value === raw)?.label ?? raw;
|
||||
}
|
||||
if (raw !== null && typeof raw === "object" && !Array.isArray(raw)) {
|
||||
const entries = Object.entries(raw as Record<string, unknown>).filter(
|
||||
(entry): entry is [string, GuardrailStreamScope] => isGuardrailStreamScope(entry[1]),
|
||||
);
|
||||
if (entries.length === 0) return "";
|
||||
return entries.map(([mode, scope]) => `${mode}: ${formatGuardrailStreamScope(scope)}`).join(", ");
|
||||
}
|
||||
return "";
|
||||
};
|
||||
|
||||
export const streamScopeForUpdate = (
|
||||
modes: string[],
|
||||
nextByMode: Record<string, GuardrailStreamScope>,
|
||||
previousRaw: unknown,
|
||||
previousModes: string[] = modes,
|
||||
): GuardrailStreamScope | Record<string, GuardrailStreamScope> | undefined => {
|
||||
const previousMap: Record<string, unknown> =
|
||||
previousRaw !== null && typeof previousRaw === "object" && !Array.isArray(previousRaw)
|
||||
? (previousRaw as Record<string, unknown>)
|
||||
: {};
|
||||
const preserved: Record<string, GuardrailStreamScope> = Object.fromEntries(
|
||||
Object.entries(previousMap).filter(
|
||||
(entry): entry is [string, GuardrailStreamScope] => !modes.includes(entry[0]) && isGuardrailStreamScope(entry[1]),
|
||||
),
|
||||
);
|
||||
const nextModes = [...modes, ...Object.keys(preserved)];
|
||||
const nextStreamScope = streamScopePayload(nextModes, {
|
||||
...preserved,
|
||||
...Object.fromEntries(modes.map((mode) => [mode, nextByMode[mode] ?? "both"])),
|
||||
});
|
||||
const previousCompareModes = Object.keys(previousMap).length > 0 ? Object.keys(previousMap) : previousModes;
|
||||
const previousStreamScope = streamScopePayload(
|
||||
previousCompareModes,
|
||||
streamScopeByModeFromConfig(previousRaw, previousCompareModes),
|
||||
);
|
||||
if (JSON.stringify(nextStreamScope ?? "both") === JSON.stringify(previousStreamScope ?? "both")) {
|
||||
return undefined;
|
||||
}
|
||||
return nextStreamScope ?? "both";
|
||||
};
|
||||
|
|
|
|||
|
|
@ -102,7 +102,7 @@ export interface PolicyAttachmentListResponse {
|
|||
|
||||
export interface PipelineStepResult {
|
||||
guardrail_name: string;
|
||||
outcome: "pass" | "fail" | "error";
|
||||
outcome: "pass" | "fail" | "error" | "skip";
|
||||
action_taken: string;
|
||||
modified_data: Record<string, any> | null;
|
||||
error_detail: string | null;
|
||||
|
|
|
|||
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -27393,6 +27393,13 @@ export interface components {
|
|||
* @default true
|
||||
*/
|
||||
sticky_session_routing: boolean | null;
|
||||
/**
|
||||
* Stream Scope
|
||||
* @description Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.
|
||||
*/
|
||||
stream_scope?: ("streaming" | "non_streaming" | "both") | {
|
||||
[key: string]: "streaming" | "non_streaming" | "both";
|
||||
} | null;
|
||||
/**
|
||||
* Template Id
|
||||
* @description The ID of your Model Armor template
|
||||
|
|
@ -37786,6 +37793,13 @@ export interface components {
|
|||
* @default true
|
||||
*/
|
||||
sticky_session_routing: boolean | null;
|
||||
/**
|
||||
* Stream Scope
|
||||
* @description Whether this guardrail runs on streaming requests, non-streaming requests, or both. A string applies to every configured mode. A map overrides named modes (pre_call, during_call, post_call, ...); omitted keys default to both. Unset means both, matching historical behavior.
|
||||
*/
|
||||
stream_scope?: ("streaming" | "non_streaming" | "both") | {
|
||||
[key: string]: "streaming" | "non_streaming" | "both";
|
||||
} | null;
|
||||
/**
|
||||
* Template Id
|
||||
* @description The ID of your Model Armor template
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue