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:
devin-ai-integration[bot] 2026-10-08 13:46:53 -07:00 • committed by GitHub
parent 7ec2a94cd8
commit 85a3869dfd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
42 changed files with 6315 additions and 367 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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

File diff suppressed because it is too large Load diff

File diff suppressed because it is too large Load diff

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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]),
);

View file

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

View file

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

View file

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

View file

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

View file

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