chore: merge litellm_internal_staging into litellm_fix_responses_bridge_streaming_contract

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mateo 2026-07-25 00:20:31 +00:00
commit f68abdc861
105 changed files with 5765 additions and 1900 deletions

View file

@ -80,7 +80,7 @@ jobs:
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'

View file

@ -55,7 +55,7 @@ jobs:
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Generate Prisma client
env:

View file

@ -64,6 +64,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
# Copy full source tree
@ -84,6 +85,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -62,6 +62,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
# Copy full source tree
@ -82,6 +83,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -68,6 +68,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
# Copy full source tree
@ -94,6 +95,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3 \
--no-sources-package litellm-proxy-extras; \
else \
@ -102,6 +104,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3; \
fi

View file

@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import (
record_output_item_chunk,
record_output_text_chunk,
)
from litellm.responses.utils import normalize_responses_api_stream_options
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionReasoningItem,
@ -320,6 +321,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
responses_api_request["tool_choice"] = ( # type: ignore[assignment]
self._normalize_tool_choice_for_responses_api(value)
)
elif key == "stream_options":
stream_options = normalize_responses_api_stream_options(value)
if stream_options is not None:
responses_api_request["stream_options"] = stream_options
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
responses_api_request[key] = value # type: ignore
elif key == "previous_response_id":
@ -360,8 +365,6 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
continue
if key == "instructions" and instructions:
request_data["instructions"] = instructions
elif key == "stream_options" and isinstance(value, dict):
request_data["stream_options"] = value.get("include_obfuscation")
elif key == "user" and isinstance(value, str):
# OpenAI API requires user param to be max 64 chars - truncate if longer
if len(value) <= 64:

View file

@ -1,3 +1,4 @@
import contextvars
import hashlib
import os
import secrets
@ -16,7 +17,11 @@ from typing import (
)
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
redact_nested_match_and_regex_keys,
)
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.secret_managers.main import str_to_bool
@ -64,6 +69,10 @@ from litellm.exceptions import (
# proxy's metadata sanitizer.
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
_guardrail_self_recorded: contextvars.ContextVar[bool] = contextvars.ContextVar(
"litellm_guardrail_self_recorded", default=False
)
def _strict_guardrail_modes_enabled() -> bool:
"""Whether guardrail-mode validation raises (default) or logs a warning.
@ -102,6 +111,8 @@ 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
records_own_guardrail_information: ClassVar[bool] = False
def __init__(
self,
guardrail_name: Optional[str] = None,
@ -117,6 +128,7 @@ class CustomGuardrail(CustomLogger):
on_sensitive_data: Optional[str] = None,
sensitive_data_route_to_model: Optional[str] = None,
sticky_session_routing: bool = True,
run_in_parallel: bool = False,
only_scan_new_messages: bool = False,
**kwargs,
):
@ -136,6 +148,9 @@ class CustomGuardrail(CustomLogger):
on_sensitive_data: Action when sensitive data is detected. 'block' (default) or 'route'
sensitive_data_route_to_model: Model to route to when on_sensitive_data='route'
sticky_session_routing: When True, all subsequent requests in the session use the same model
run_in_parallel: When True, this pre_call or post_call guardrail runs concurrently with
other opted-in guardrails of the same hook. Only safe for block-only guardrails that
do not mutate the request or response.
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -150,6 +165,7 @@ class CustomGuardrail(CustomLogger):
self.on_sensitive_data: Optional[str] = on_sensitive_data
self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model
self.sticky_session_routing: bool = sticky_session_routing
self.run_in_parallel: bool = run_in_parallel
self.only_scan_new_messages: bool = only_scan_new_messages
if supported_event_hooks:
@ -944,17 +960,10 @@ class CustomGuardrail(CustomLogger):
# should not happen
container[key] = [existing, slg]
if "metadata" in request_data:
if request_data["metadata"] is None:
request_data["metadata"] = {}
_append_guardrail_info(request_data["metadata"])
elif "litellm_metadata" in request_data:
_append_guardrail_info(request_data["litellm_metadata"])
else:
# Ensure guardrail info is always logged (e.g. proxy may not have set
# metadata yet). Attach to "metadata" so spend log / standard logging see it.
request_data["metadata"] = {}
_append_guardrail_info(request_data["metadata"])
_, metadata_bucket = get_or_create_metadata_bucket(request_data)
_append_guardrail_info(metadata_bucket)
_guardrail_self_recorded.set(True)
# Emit the otel guardrail span here, where every guardrail execution lands,
# rather than relying on a post-call hook that does not fire on every path
@ -1211,7 +1220,7 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object)
"""
if logging_obj is None:
return
meta_src = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
meta_src = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
slg_info = meta_src.get("standard_logging_guardrail_information")
if not slg_info:
return
@ -1238,8 +1247,20 @@ def log_guardrail_information(func):
(structured detections, tracing detail) than this decorator's
"allow"/"mask"/raw-response default. To avoid double-recording in that
case (which would emit two spans, two Datadog records, two spend-log
entries, etc.), snapshot the entry count before invocation: if the
wrapped function already appended its own entry, skip the auto-record.
entries, etc.), a context-local flag records whether the wrapped function
appended its own entry; if so, the auto-record is skipped. The flag is a
``ContextVar`` rather than a count of entries in the shared ``request_data``
so it stays correct when guardrails run concurrently (asyncio copies the
context into each gathered task): counting shared entries would let one
guardrail's append hide another guardrail's missing record.
A guardrail that only records an entry when it actually runs (e.g.
``HeadroomGuardrail``, which returns the inputs untouched on an endpoint
whose payload it cannot act on) sets ``records_own_guardrail_information =
True`` so the auto-record is skipped even on the return paths where it
recorded nothing; otherwise a no-op early return would be logged as an
"allow"/"success" run even though the guardrail did nothing. The exception
branch below still records so a genuine failure is not lost.
"""
import functools
import inspect
@ -1259,16 +1280,6 @@ def log_guardrail_information(func):
return GuardrailEventHooks.post_call
return None
def _count_recorded_guardrail_entries(request_data: dict) -> int:
total = 0
for container_key in ("metadata", "litellm_metadata"):
container = request_data.get(container_key)
if isinstance(container, dict):
entries = container.get("standard_logging_guardrail_information")
if isinstance(entries, list):
total += len(entries)
return total
@functools.wraps(func)
async def async_wrapper(*args, **kwargs):
start_time = datetime.now() # Move start_time inside the wrapper
@ -1282,10 +1293,10 @@ def log_guardrail_information(func):
original_inputs = kwargs.get("inputs")
logging_obj = kwargs.get("logging_obj")
entries_before = _count_recorded_guardrail_entries(request_data)
self_recorded_token = _guardrail_self_recorded.set(False)
try:
response = await func(*args, **kwargs)
if _count_recorded_guardrail_entries(request_data) > entries_before:
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
return response
return self._process_response(
response=response,
@ -1297,7 +1308,7 @@ def log_guardrail_information(func):
original_inputs=original_inputs,
)
except Exception as e:
if _count_recorded_guardrail_entries(request_data) > entries_before:
if _guardrail_self_recorded.get():
raise
return self._process_error(
e=e,
@ -1308,6 +1319,7 @@ def log_guardrail_information(func):
event_type=event_type,
)
finally:
_guardrail_self_recorded.reset(self_recorded_token)
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
@functools.wraps(func)
@ -1323,10 +1335,10 @@ def log_guardrail_information(func):
original_inputs = kwargs.get("inputs")
logging_obj = kwargs.get("logging_obj")
entries_before = _count_recorded_guardrail_entries(request_data)
self_recorded_token = _guardrail_self_recorded.set(False)
try:
response = func(*args, **kwargs)
if _count_recorded_guardrail_entries(request_data) > entries_before:
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
return response
return self._process_response(
response=response,
@ -1336,7 +1348,7 @@ def log_guardrail_information(func):
original_inputs=original_inputs,
)
except Exception as e:
if _count_recorded_guardrail_entries(request_data) > entries_before:
if _guardrail_self_recorded.get():
raise
return self._process_error(
e=e,
@ -1345,6 +1357,7 @@ def log_guardrail_information(func):
event_type=event_type,
)
finally:
_guardrail_self_recorded.reset(self_recorded_token)
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
@functools.wraps(func)

View file

@ -883,8 +883,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
request_data: dict,
parent_span: Optional[Any],
) -> None:
"""Emit ``guardrail`` spans from ``request_data["metadata"]
["standard_logging_guardrail_information"]``.
"""Emit ``guardrail`` spans from the request's proxy-internal metadata bucket
(``standard_logging_guardrail_information``).
Routed through ``_create_guardrail_span`` so the dedupe state in
``_otel_internal`` is honoured — if ``_handle_failure`` already
@ -892,7 +892,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"""
from opentelemetry import trace as _trace
metadata = (request_data or {}).get("metadata") or {}
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
request_data = request_data or {}
metadata = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
guardrail_information = metadata.get("standard_logging_guardrail_information")
if not guardrail_information:
return

View file

@ -18,12 +18,14 @@ before the LLM call even starts), so a guardrail is a sibling of the LLM call,
not a child of it. The emitter parents every span to the ambient OTel context
(the active server span), which matches this.
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) are intentionally NOT in this
tree. Per the OTel GenAI MCP semconv, MCP and the HTTP transport are independent
contexts, so an MCP span parents to the trace context the client propagated in
``params._meta`` (or starts its own root when none is propagated) and records the
``PROXY_REQUEST`` transport span as a span *link*, never a parent. The registry
encodes this as ``parent=None, links=PROXY_REQUEST``.
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) have two shapes, chosen at emit
time by :func:`resolve_mcp_span_context`. When the client propagates trace context
in ``params._meta`` MCP and the HTTP transport are independent contexts per the
OTel GenAI MCP semconv, so the span parents to that propagated context and records
the ``PROXY_REQUEST`` transport span as a span *link*, never a parent — the shape
this registry's ``parent=None, links=PROXY_REQUEST`` entry encodes. When nothing is
propagated (the common case) the span nests under the transport span of the request
carrying that message, so the tool call stays in one trace.
Not every service call becomes a span — :func:`span_role_for_service` decides:
@ -89,12 +91,13 @@ class SpanSpec:
SPAN_REGISTRY: dict[SpanRole, SpanSpec] = {
SpanRole.PROXY_REQUEST: SpanSpec(SpanRole.PROXY_REQUEST, LiteLLMSpanKind.SERVER, parent=None),
SpanRole.LLM_CALL: SpanSpec(SpanRole.LLM_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
# MCP and the HTTP transport are independent contexts (OTel GenAI MCP semconv),
# so an MCP span does not nest under the transport span. The proxy is an MCP
# client to the upstream server, so it's a CLIENT span; it parents to the trace
# context the client propagated in ``params._meta`` (or starts its own root when
# none is propagated) and records the PROXY_REQUEST transport span as a span
# *link*, never a parent — hence ``parent=None, links=PROXY_REQUEST``.
# The proxy is an MCP client to the upstream server, so MCP spans are CLIENT
# spans. With trace context propagated in ``params._meta``, MCP and the HTTP
# transport are independent contexts (OTel GenAI MCP semconv): the span parents
# to the propagated context and records the PROXY_REQUEST transport span as a
# span *link*, never a parent — the shape ``parent=None, links=PROXY_REQUEST``
# encodes. With nothing propagated, ``resolve_mcp_span_context`` nests the span
# under that message's transport span instead, keeping the call in one trace.
SpanRole.MCP_TOOL_CALL: SpanSpec(
SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=None, links=SpanRole.PROXY_REQUEST
),

View file

@ -5,7 +5,14 @@ from typing import Mapping
from opentelemetry import baggage
from opentelemetry.context import Context, get_current
from opentelemetry.trace import Link, Span, get_current_span, set_span_in_context
from opentelemetry.trace import (
Link,
NonRecordingSpan,
Span,
SpanContext,
get_current_span,
set_span_in_context,
)
from opentelemetry.trace.propagation.tracecontext import (
TraceContextTextMapPropagator,
)
@ -72,6 +79,62 @@ def reset_mcp_message_trace_carrier(token: "Token[Mapping[str, str] | None]") ->
_mcp_message_trace_carrier.reset(token)
# The transport span of the HTTP request carrying the CURRENT MCP message, as a
# plain ``SpanContext`` so it can cross a task boundary.
#
# ``_request_root_span`` above cannot be used for MCP: a *stateful* streamable-HTTP
# session runs every message on the single task spawned by that session's
# ``initialize`` POST, so the ContextVar the ASGI request task writes at auth time
# is frozen at ``initialize`` there and never sees the later ``tools/call`` POSTs.
# Reading it from the message handler would parent every tool call in the session
# to the first request's (already ended) server span. The gateway instead resolves
# the current message's transport span on the request task and hands it over the
# same way it hands over per-request auth, and the handler publishes it here for
# the span emitter to pick up.
_mcp_message_transport_span_context: "ContextVar[SpanContext | None]" = ContextVar(
"litellm_otel_mcp_message_transport_span_context", default=None
)
def set_mcp_message_transport_span_context(
span_context: "SpanContext | None",
) -> "Token[SpanContext | None]":
"""Publish the transport span of the request carrying the current MCP message.
Returns the reset token; the caller must reset it once the message is handled
so the transport never leaks to the next message on the same session task.
"""
return _mcp_message_transport_span_context.set(span_context)
def reset_mcp_message_transport_span_context(token: "Token[SpanContext | None]") -> None:
_mcp_message_transport_span_context.reset(token)
def request_root_span_context() -> "SpanContext | None":
"""The anchored request root span's context, safe to hand to another task.
A ``SpanContext`` is an immutable value, unlike the live ``Span``, so passing it
across the MCP session-task boundary cannot keep a finished span alive or invite
writes to it from the wrong request.
"""
span = request_root_span()
return span.get_span_context() if span is not None else None
def _mcp_transport_span_context() -> "SpanContext | None":
"""The transport span an MCP message span should attach to.
Prefers the transport the gateway published for this specific message; falls
back to the ambient request anchor for paths that emit an MCP span on the
request task itself (the REST MCP endpoints, the SDK).
"""
published = _mcp_message_transport_span_context.get()
if published is not None and published.is_valid:
return published
return request_root_span_context()
def set_request_baggage(values: Mapping[str, str], context: Context | None = None) -> Context:
"""Return a context with ``values`` written into Baggage."""
ctx = context
@ -132,33 +195,44 @@ def resolve_request_span_context() -> Context:
def resolve_mcp_span_context(
carrier: "Mapping[str, str] | None" = None,
) -> "tuple[Context, tuple[Link, ...]]":
"""Parent context + links for an MCP message span, per the OTel GenAI MCP semconv.
"""Parent context + links for an MCP message span.
MCP and the underlying transport (HTTP) are independent lifecycles — one
streamable-HTTP session multiplexes many messages, so nesting the message span
under the HTTP/session span is wrong (it renders the message at the session's
start, skewed by however long the session has been open). Instead:
When the client propagates W3C trace context in the request's ``params._meta``
(SEP-414), MCP and the underlying transport are independent lifecycles — one
streamable-HTTP session multiplexes many messages, and the client's own span is
the truthful parent. So, per the OTel GenAI MCP semconv:
* parent to the trace context the client propagated in the request's
``params._meta`` (a *remote* parent), and
* record the transport/session span as a *link*, never the parent.
* parent to the trace context the client propagated (a *remote* parent), and
* record the transport span as a *link*, never the parent.
Almost no client implements SEP-414 yet, so in practice nothing is propagated.
Rooting the span there splits a single tool call into two disconnected traces
joined only by a link, which is how it surfaces in APM: the ``POST`` transaction
and the ``tools/call`` span share no trace. With no remote parent to honor,
parent to the transport span of the request carrying this message instead, so
the call stays in one trace; no link is added since the transport is now the
real parent. The transport comes from :func:`_mcp_transport_span_context`, which
is the *current message's* POST rather than whatever request happened to open
the session, so a long-lived session does not glue every message under its
first request. With neither a remote parent nor a transport the returned context
carries no span and the span legitimately starts its own root trace.
Only trace context (``traceparent``/``tracestate``) is extracted, never the
client's W3C Baggage: ``params._meta`` is caller-controlled, and the otel
baggage processor stamps allowlisted baggage keys (``litellm.team.id``,
``litellm.metadata.*``, ...) onto the span as attributes, so honoring remote
baggage would let a client spoof a span's identity attribution.
With no propagated context the returned context carries no span, so the span
starts its own root trace (still linked to the transport). The base context is
explicitly empty so an absent ``traceparent`` can never fall through to the
ambient (stale session) span.
baggage would let a client spoof a span's identity attribution. The base context
for extraction is explicitly empty so an absent or malformed ``traceparent`` can
never fall through to the ambient (stale session) span.
"""
source = carrier if carrier is not None else _mcp_message_trace_carrier.get()
parent = _PROPAGATOR.extract(dict(source or {}), context=Context())
transport = request_root_span()
links = (Link(transport.get_span_context()),) if transport is not None else ()
return parent, links
transport = _mcp_transport_span_context()
if is_recordable_span(get_current_span(parent)):
return parent, (Link(transport),) if transport is not None else ()
if transport is not None:
return context_from_span(NonRecordingSpan(transport)), ()
return parent, ()
def is_recordable_span(obj: object) -> bool:

View file

@ -195,6 +195,25 @@ def get_metadata_variable_name_from_kwargs(
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
def get_or_create_metadata_bucket(
request_data: dict,
) -> tuple[Literal["metadata", "litellm_metadata"], dict]:
"""
Return the proxy-internal metadata bucket for this request, creating it if absent.
Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI
``metadata`` field can remain provider-safe (string values only). Every writer and
reader of proxy-internal metadata resolves the bucket through here, so a caller that
supplies its own ``metadata`` field cannot split them across two dicts.
"""
metadata_key = get_metadata_variable_name_from_kwargs(request_data)
metadata_bucket = request_data.get(metadata_key)
if not isinstance(metadata_bucket, dict):
metadata_bucket = {}
request_data[metadata_key] = metadata_bucket
return metadata_key, metadata_bucket
def get_litellm_metadata_from_kwargs(kwargs: dict):
"""
Helper to get litellm metadata from all litellm request kwargs

View file

@ -600,9 +600,15 @@ class AnthropicMessagesHandler(BaseTranslation):
guardrail_inputs["tool_calls"] = tool_calls_list
try:
prepared_request_data = self._prepare_request_data(
request_data,
model_response,
user_api_key_dict,
key="response",
)
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=guardrail_inputs,
request_data=request_data if request_data is not None else {},
request_data=prepared_request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)
@ -618,9 +624,15 @@ class AnthropicMessagesHandler(BaseTranslation):
string_so_far = self.get_streaming_string_so_far(responses_so_far)
try:
prepared_request_data = self._prepare_request_data(
request_data,
responses_so_far,
user_api_key_dict,
key="responses",
)
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [string_so_far]},
request_data=request_data if request_data is not None else {},
request_data=prepared_request_data,
input_type="response",
logging_obj=litellm_logging_obj,
)

View file

@ -123,7 +123,8 @@ class VertexAIBatchTransformation:
Gets the output file id from the Vertex AI Batch response
"""
output_file_id: str = response.get("outputInfo", OutputInfo()).get("gcsOutputDirectory", "")
output_info = response.get("outputInfo") or OutputInfo()
output_file_id: str = output_info.get("gcsOutputDirectory", "")
if output_file_id:
output_file_id = output_file_id.rstrip("/") + "/predictions.jsonl"
if output_file_id and output_file_id != "/predictions.jsonl":

View file

@ -1,9 +1,12 @@
from typing import Dict, List, Optional
from typing import TYPE_CHECKING, Dict, List, Optional
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
from litellm.proxy._types import UserAPIKeyAuth
if TYPE_CHECKING:
from opentelemetry.trace import SpanContext
class MCPAuthenticatedUser(AuthenticatedUser):
"""
@ -16,6 +19,8 @@ class MCPAuthenticatedUser(AuthenticatedUser):
4. Server-specific authentication headers
5. OAuth2 headers
6. Raw headers - allows forwarding specific headers to the MCP server, specified by the admin.
7. Transport span context - the tracing span of the HTTP request carrying the current
message, which a stateful session's message handler cannot read from its own task.
"""
def __init__(
@ -28,6 +33,7 @@ class MCPAuthenticatedUser(AuthenticatedUser):
mcp_protocol_version: Optional[str] = None,
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
transport_span_context: Optional["SpanContext"] = None,
):
self.user_api_key_auth = user_api_key_auth
self.mcp_auth_header = mcp_auth_header
@ -37,3 +43,4 @@ class MCPAuthenticatedUser(AuthenticatedUser):
self.oauth2_headers = oauth2_headers
self.raw_headers = raw_headers
self.client_ip = client_ip
self.transport_span_context = transport_span_context

View file

@ -15,6 +15,7 @@ import types
import uuid
from datetime import datetime
from typing import (
TYPE_CHECKING,
Any,
AsyncIterator,
Callable,
@ -107,6 +108,9 @@ _MAX_STATEFUL_SESSIONS_PER_OWNER = 100
# arbitrarily large body just to make a routing decision.
_MCP_ROUTING_PEEK_MAX_BYTES = 4096
if TYPE_CHECKING:
from opentelemetry.trace import SpanContext
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
"""Remove a (user_id, server_id) entry from the BYOK credential cache.
@ -242,10 +246,12 @@ def _mcp_meta_trace_carrier(req_ctx: object) -> Optional[dict[str, str]]:
"""The W3C trace context (``traceparent``/``tracestate``) the MCP client
propagated in the request's ``params._meta`` (SEP-414), or ``None``.
Per the OTel MCP semconv the MCP span parents to this propagated context rather
than to the HTTP/session transport (which is recorded as a link instead), so a
streamable-HTTP session that multiplexes many messages does not glue every
message under the session's first request. The client's W3C Baggage is
When present, per the OTel MCP semconv the MCP span parents to this propagated
context rather than to the HTTP transport (which is recorded as a link instead).
When absent, the span nests under the transport span of the request carrying
this specific message, so a streamable-HTTP session that multiplexes many
messages still does not glue every message under the session's first request;
see ``resolve_mcp_span_context``. The client's W3C Baggage is
deliberately excluded: it is caller-controlled, and the otel baggage processor
stamps allowlisted baggage keys (``litellm.team.id``, ``litellm.metadata.*``,
...) onto the span, so honoring remote baggage would let a client spoof a
@ -288,6 +294,56 @@ def _otel_reset_mcp_trace_carrier(token: object) -> None:
return
def _otel_request_transport_span_context() -> Optional["SpanContext"]:
"""The tracing span of the HTTP request being handled, as a portable value.
Resolved on the ASGI request task, where the proxy's server span is anchored,
and carried to the MCP message handler on the authenticated-user object. A
stateful streamable-HTTP session handles every message on the task spawned by
its ``initialize`` POST, so the handler's own task cannot see later requests'
spans; this is the same reason per-request auth is carried across rather than
read from a ContextVar. Lazily imported so opentelemetry stays an optional
dependency; returns ``None`` when otel_v2 is unavailable or no request span is
anchored."""
try:
from litellm.integrations.otel.plumbing.context import (
request_root_span_context,
)
return request_root_span_context()
except ImportError:
return None
def _otel_set_mcp_transport_span_context(span_context: Optional["SpanContext"]) -> object:
"""Publish the current message's transport span for the otel_v2 MCP span and
return a reset token, or ``None`` when otel_v2 is unavailable."""
if span_context is None:
return None
try:
from litellm.integrations.otel.plumbing.context import (
set_mcp_message_transport_span_context,
)
return set_mcp_message_transport_span_context(span_context)
except ImportError:
return None
def _otel_reset_mcp_transport_span_context(token: object) -> None:
"""Paired with ``_otel_set_mcp_transport_span_context``."""
if token is None:
return
try:
from litellm.integrations.otel.plumbing.context import (
reset_mcp_message_transport_span_context,
)
reset_mcp_message_transport_span_context(token)
except ImportError:
return
def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
"""Map a ``ProxyException`` to an ``HTTPException`` that preserves its real
status code and headers.
@ -654,6 +710,18 @@ if MCP_AVAILABLE:
############### MCP Server Routes #######################
########################################################
def _current_transport_span_context() -> Optional["SpanContext"]:
"""The transport span of the HTTP request carrying the message being handled.
Published by the ASGI request task onto the authenticated-user object, because
a stateful session's message handler runs on the task spawned by that session's
``initialize`` POST and so cannot read later requests' spans from its own task.
"""
auth_user = auth_context_var.get()
if not isinstance(auth_user, MCPAuthenticatedUser):
auth_user = _recover_auth_from_session()
return auth_user.transport_span_context if auth_user is not None else None
@server.list_tools()
async def handle_list_tools() -> "ListToolsResult | List[Tool]":
"""
@ -670,9 +738,11 @@ if MCP_AVAILABLE:
if req_ctx:
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
_trace_token = None
_transport_token = None
try:
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
_transport_token = _otel_set_mcp_transport_span_context(_current_transport_span_context())
# Get user authentication from context variable
(
user_api_key_auth,
@ -728,6 +798,7 @@ if MCP_AVAILABLE:
# This prevents the HTTP stream from failing and allows the client to get a response
return []
finally:
_otel_reset_mcp_transport_span_context(_transport_token)
_otel_reset_mcp_trace_carrier(_trace_token)
if _session_reset_token is not None:
active_mcp_session_var.reset(_session_reset_token)
@ -901,9 +972,11 @@ if MCP_AVAILABLE:
if req_ctx:
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
_trace_token = None
_transport_token = None
try:
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
_transport_token = _otel_set_mcp_transport_span_context(_current_transport_span_context())
# Validate arguments
(
user_api_key_auth,
@ -1042,6 +1115,7 @@ if MCP_AVAILABLE:
return response
finally:
_otel_reset_mcp_transport_span_context(_transport_token)
_otel_reset_mcp_trace_carrier(_trace_token)
if _session_reset_token is not None:
active_mcp_session_var.reset(_session_reset_token)
@ -4197,6 +4271,7 @@ if MCP_AVAILABLE:
session_id=session_id if use_stateful else None,
touch_last_seen=(scope.get("method") or "").upper() != "DELETE",
copy_existing_session_auth_context=is_initialize,
transport_span_context=_otel_request_transport_span_context(),
)
local_send = send
if use_stateful and is_initialize:
@ -4421,6 +4496,7 @@ if MCP_AVAILABLE:
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
transport_span_context: Optional["SpanContext"] = None,
) -> None:
auth_user.user_api_key_auth = user_api_key_auth
auth_user.mcp_auth_header = mcp_auth_header
@ -4429,6 +4505,7 @@ if MCP_AVAILABLE:
auth_user.oauth2_headers = oauth2_headers
auth_user.raw_headers = raw_headers
auth_user.client_ip = client_ip
auth_user.transport_span_context = transport_span_context
def set_auth_context(
user_api_key_auth: Optional[UserAPIKeyAuth],
@ -4438,6 +4515,7 @@ if MCP_AVAILABLE:
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
client_ip: Optional[str] = None,
transport_span_context: Optional["SpanContext"] = None,
) -> MCPAuthenticatedUser:
"""
Set the UserAPIKeyAuth in the auth context variable.
@ -4448,6 +4526,7 @@ if MCP_AVAILABLE:
mcp_servers: Optional list of server names and access groups to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
client_ip: Client IP address for MCP access control
transport_span_context: Tracing span of the HTTP request carrying this message
"""
auth_user = MCPAuthenticatedUser(
user_api_key_auth=user_api_key_auth,
@ -4457,6 +4536,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
transport_span_context=transport_span_context,
)
auth_context_var.set(auth_user)
return auth_user
@ -4472,6 +4552,7 @@ if MCP_AVAILABLE:
session_id: Optional[str] = None,
touch_last_seen: bool = True,
copy_existing_session_auth_context: bool = False,
transport_span_context: Optional["SpanContext"] = None,
) -> MCPAuthenticatedUser:
auth_user = _stateful_session_auth_contexts.get(session_id) if session_id else None
if auth_user is not None and session_id is not None:
@ -4486,6 +4567,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
transport_span_context=transport_span_context,
)
_update_auth_context(
auth_user=auth_user,
@ -4496,6 +4578,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
transport_span_context=transport_span_context,
)
auth_context_var.set(auth_user)
return auth_user
@ -4507,6 +4590,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=client_ip,
transport_span_context=transport_span_context,
)
def _wrap_send_with_stateful_session_auth_context(

View file

@ -1,12 +1,16 @@
import copy
import os
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional
import litellm
from litellm import get_secret
from litellm._logging import verbose_proxy_logger
from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
)
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.proxy._types import CommonProxyErrors, LiteLLMPromptInjectionParams
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
@ -406,23 +410,6 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]:
return headers
def get_metadata_variable_name_from_kwargs(
kwargs: dict,
) -> Literal["metadata", "litellm_metadata"]:
"""
Helper to return what the "metadata" field should be called in the request data
- New endpoints return `litellm_metadata`
- Old endpoints return `metadata`
Context:
- LiteLLM used `metadata` as an internal field for storing metadata
- OpenAI then started using this field for their metadata
- LiteLLM is now moving to using `litellm_metadata` for our metadata
"""
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset(
{
"applied_policies",
@ -450,23 +437,6 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS = frozenset(
)
def _get_or_create_proxy_metadata_bucket(
request_data: Dict,
) -> tuple[Literal["metadata", "litellm_metadata"], dict]:
"""
Return the proxy-internal metadata bucket for this request.
Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI
``metadata`` field can remain provider-safe (string values only).
"""
metadata_key = get_metadata_variable_name_from_kwargs(request_data)
metadata_bucket = request_data.get(metadata_key)
if not isinstance(metadata_bucket, dict):
metadata_bucket = {}
request_data[metadata_key] = metadata_bucket
return metadata_key, metadata_bucket
def sanitize_openai_provider_metadata(
metadata: Optional[Dict[str, Any]],
) -> Optional[Dict[str, str]]:
@ -496,7 +466,7 @@ def sanitize_openai_provider_metadata(
def add_guardrail_to_applied_guardrails_header(request_data: Dict, guardrail_name: Optional[str]):
if guardrail_name is None:
return
_, _metadata = _get_or_create_proxy_metadata_bucket(request_data)
_, _metadata = get_or_create_metadata_bucket(request_data)
if "applied_guardrails" in _metadata:
if guardrail_name not in _metadata["applied_guardrails"]:
_metadata["applied_guardrails"].append(guardrail_name)
@ -513,7 +483,7 @@ def add_policy_to_applied_policies_header(request_data: Dict, policy_name: Optio
"""
if policy_name is None:
return
_, _metadata = _get_or_create_proxy_metadata_bucket(request_data)
_, _metadata = get_or_create_metadata_bucket(request_data)
if "applied_policies" in _metadata:
if policy_name not in _metadata["applied_policies"]:
_metadata["applied_policies"].append(policy_name)
@ -531,7 +501,7 @@ def add_policy_sources_to_metadata(request_data: Dict, policy_sources: Dict[str,
"""
if not policy_sources:
return
_, _metadata = _get_or_create_proxy_metadata_bucket(request_data)
_, _metadata = get_or_create_metadata_bucket(request_data)
existing = _metadata.get("policy_sources", {})
if not isinstance(existing, dict):
existing = {}

View file

@ -36,6 +36,10 @@ SSO_DESCRIPTORS: tuple[FieldDescriptor, ...] = (
FieldDescriptor("generic_token_endpoint", "generic_token_endpoint", "GENERIC_TOKEN_ENDPOINT"),
FieldDescriptor("generic_userinfo_endpoint", "generic_userinfo_endpoint", "GENERIC_USERINFO_ENDPOINT"),
FieldDescriptor("generic_scope", "generic_scope", "GENERIC_SCOPE", default="openid email profile"),
FieldDescriptor("saml_idp_metadata_url", "saml_idp_metadata_url", "SAML_IDP_METADATA_URL"),
FieldDescriptor("saml_idp_metadata_xml", "saml_idp_metadata_xml", "SAML_IDP_METADATA_XML"),
FieldDescriptor("saml_sp_entity_id", "saml_sp_entity_id", "SAML_SP_ENTITY_ID"),
FieldDescriptor("saml_allow_unsolicited", "saml_allow_unsolicited", "SAML_ALLOW_UNSOLICITED"),
FieldDescriptor("proxy_base_url", "proxy_base_url", "PROXY_BASE_URL"),
)

View file

@ -4,7 +4,7 @@ import json
import re
import time
import uuid
from typing import TYPE_CHECKING, Any, List, Literal, Optional
from typing import TYPE_CHECKING, Any, ClassVar, List, Literal, Optional
import httpx
from fastapi import HTTPException
@ -209,6 +209,8 @@ def _build_responses_followup_items(
class HeadroomGuardrail(CustomGuardrail):
records_own_guardrail_information: ClassVar[bool] = True
@classmethod
def get_supported_event_hooks(cls) -> List[GuardrailEventHooks]:
return [
@ -481,7 +483,21 @@ class HeadroomGuardrail(CustomGuardrail):
)
end_time = time.time()
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
if not compression_succeeded:
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response={"error": "headroom compression unavailable; request forwarded uncompressed"},
request_data=request_data,
guardrail_status="guardrail_failed_to_respond",
guardrail_provider=HEADROOM_GUARDRAIL_PROVIDER,
start_time=start_time,
end_time=end_time,
duration=end_time - start_time,
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType]
self.add_standard_logging_guardrail_information_to_request_data(
@ -493,6 +509,7 @@ class HeadroomGuardrail(CustomGuardrail):
end_time=end_time,
duration=end_time - start_time,
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
hashes = extract_hashes_from_messages(compressed)
if not hashes:

View file

@ -30,6 +30,10 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import (
@ -432,7 +436,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
Override to store only the Model Armor API response, not the entire data dict.
This prevents circular references in logging.
"""
metadata = (request_data.get("metadata") or {}) if isinstance(request_data, dict) else {}
metadata = (
request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
if isinstance(request_data, dict)
else {}
)
guardrail_response = metadata.get("_model_armor_response", {})
# Determine status – default to "success" but prefer the explicit value if present.
@ -471,7 +479,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
blocking, while fail_on_error still governs real Model Armor API errors.
"""
from litellm.proxy.common_utils.callback_utils import (
_get_or_create_proxy_metadata_bucket,
add_guardrail_to_applied_guardrails_header,
)
@ -491,7 +498,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
# Use the same metadata bucket the header helper writes to, so the logged Model Armor
# payload and status land where _process_response reads them on every route.
_, metadata = _get_or_create_proxy_metadata_bucket(data)
_, metadata = get_or_create_metadata_bucket(data)
fail_on_error = bool(self.optional_params.get("fail_on_error", True))
if unscannable_references > 0:
@ -607,7 +614,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
# overwritten by another coroutine.
blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content)
if isinstance(data, dict):
metadata = data.setdefault("metadata", {}) # ensures metadata exists and is unique per request
_, metadata = get_or_create_metadata_bucket(data) # ensures metadata exists and is unique per request
# Accumulate so a prior file scan on the same request is not overwritten by this text scan.
metadata["_model_armor_response"] = self._append_armor_response(
metadata.get("_model_armor_response"),
@ -702,7 +709,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
blocked = self._should_block_content(armor_response, allow_sanitization=self.mask_request_content)
# Store the armor response for logging
if isinstance(data, dict):
metadata = data.setdefault("metadata", {})
_, metadata = get_or_create_metadata_bucket(data)
# Accumulate so a prior file scan on the same request is not overwritten by this text scan.
metadata["_model_armor_response"] = self._append_armor_response(
metadata.get("_model_armor_response"),
@ -868,7 +875,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
# Attach Model Armor response & status to this request's metadata to avoid race conditions
if isinstance(request_data, dict):
metadata = request_data.setdefault("metadata", {})
_, metadata = get_or_create_metadata_bucket(request_data)
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
metadata["_model_armor_status"] = (
"blocked" if self._should_block_content(armor_response) else "success"

View file

@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Literal, NoReturn
from urllib.parse import urlsplit
import httpx
from pydantic import ValidationError
from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
@ -24,11 +24,12 @@ from litellm.integrations.custom_guardrail import (
log_guardrail_information,
)
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
STRAIKER_WEBHOOK_SCHEMA_VERSION,
StraikerGuardrailConfigModel,
@ -42,7 +43,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
StraikerWebhookStream,
StraikerWebhookUsage,
)
from litellm.types.utils import GenericGuardrailAPIInputs, Usage
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -57,6 +58,7 @@ RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504})
UNREACHABLE_STATUS = frozenset({502, 503, 504})
_APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"})
_OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool)
_JSON_DICT_ADAPTER = TypeAdapter(dict[str, object])
@dataclass(frozen=True, slots=True)
@ -137,6 +139,44 @@ def _resolve_destination(request_data: dict) -> str | None:
return None
def _route_has_translation(request_data: dict) -> bool:
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.llms import load_guardrail_translation_mappings
route = _as_dict(request_data.get("litellm_metadata")).get("user_api_key_request_route")
if not isinstance(route, str) or not route:
return False
mappings = load_guardrail_translation_mappings()
return any(call_type in mappings for call_type in get_call_types_for_route(route) or ())
def _request_structured_messages(request_data: dict) -> list[dict[str, Any]] | None:
messages = request_data.get("messages")
if messages:
return messages if isinstance(messages, list) else None
if not _route_has_translation(request_data):
return None
return resolve_structured_messages(messages=None, request_kwargs=request_data)
def _hook_name(value: object) -> str:
return value.value if isinstance(value, GuardrailEventHooks) else str(value)
def _configured_modes(event_hook: object) -> list[str] | None:
if isinstance(event_hook, list):
names = [_hook_name(v) for v in event_hook]
elif isinstance(event_hook, (str, GuardrailEventHooks)):
names = [_hook_name(event_hook)]
elif isinstance(event_hook, Mode):
default = event_hook.default if isinstance(event_hook.default, list) else [event_hook.default]
tags = [v for value in event_hook.tags.values() for v in (value if isinstance(value, list) else [value])]
names = [_hook_name(v) for v in (*default, *tags) if v is not None]
else:
return None
return list(dict.fromkeys(names)) or None
def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str:
call_type = (
(getattr(logging_obj, "call_type", None) if logging_obj is not None else None)
@ -146,23 +186,76 @@ def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: d
return call_type if isinstance(call_type, str) and call_type else "unknown"
def _jsonable_dict(value: object) -> dict[str, object] | None:
if isinstance(value, BaseModel):
return _JSON_DICT_ADAPTER.validate_python(value.model_dump(mode="json", exclude_none=True))
if isinstance(value, dict):
return _JSON_DICT_ADAPTER.validate_python(value)
return None
def _opaque_dict_list(value: object) -> list[dict[str, object]] | None:
if not isinstance(value, list):
return None
items = tuple(plain for item in value if (plain := _jsonable_dict(item)) is not None)
return list(items) if items else None
def _choice_terminal_reason(choice: object) -> str | None:
if isinstance(choice, dict):
return _as_optional_str(choice.get("finish_reason")) or _as_optional_str(choice.get("stop_reason"))
return _as_optional_str(getattr(choice, "finish_reason", None)) or _as_optional_str(
getattr(choice, "stop_reason", None)
)
def _response_finish_reason(response: Any) -> str | None:
if response is None:
return None
if isinstance(response, dict):
top = _as_optional_str(response.get("finish_reason")) or _as_optional_str(response.get("stop_reason"))
if top:
return top
choices = response.get("choices")
if not isinstance(choices, list):
return None
for choice in choices:
reason = _choice_terminal_reason(choice)
if reason:
return reason
return None
top = _as_optional_str(getattr(response, "finish_reason", None)) or _as_optional_str(
getattr(response, "stop_reason", None)
)
if top:
return top
choices = getattr(response, "choices", None)
if not isinstance(choices, list):
return None
for choice in choices:
reason = getattr(choice, "finish_reason", None)
if isinstance(reason, str) and reason:
reason = _choice_terminal_reason(choice)
if reason:
return reason
return None
def _as_optional_int(value: object) -> int | None:
return value if isinstance(value, int) and not isinstance(value, bool) else None
def _usage_token_count(usage: object, openai_key: str, anthropic_key: str) -> int | None:
get = usage.get if isinstance(usage, dict) else lambda key: getattr(usage, key, None)
openai_count = _as_optional_int(get(openai_key))
return openai_count if openai_count is not None else _as_optional_int(get(anthropic_key))
def _build_usage(response: object) -> StraikerWebhookUsage | None:
usage = getattr(response, "usage", None)
if not isinstance(usage, Usage):
usage = response.get("usage") if isinstance(response, dict) else getattr(response, "usage", None)
if usage is None:
return None
input_tokens = usage.prompt_tokens
output_tokens = usage.completion_tokens
input_tokens = _usage_token_count(usage, "prompt_tokens", "input_tokens")
output_tokens = _usage_token_count(usage, "completion_tokens", "output_tokens")
if input_tokens is None and output_tokens is None:
return None
return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens)
@ -234,6 +327,8 @@ class StraikerGuardrail(CustomGuardrail):
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
self.configured_modes = _configured_modes(self.event_hook)
def _webhook_url(self) -> str:
return f"{self.api_base}{WEBHOOK_PATH}"
@ -263,6 +358,7 @@ class StraikerGuardrail(CustomGuardrail):
) -> StraikerWebhookContext:
return StraikerWebhookContext(
call_surface=_resolve_call_surface(logging_obj, request_data),
mode=self.configured_modes,
model=model,
model_provider=_resolve_provider(request_data, model),
destination=_resolve_destination(request_data),
@ -287,9 +383,9 @@ class StraikerGuardrail(CustomGuardrail):
content = StraikerWebhookContent(
texts=list(inputs.get("texts") or []),
images=list(inputs.get("images") or []),
structured_messages=inputs.get("structured_messages"),
tools=inputs.get("tools"),
tool_calls=inputs.get("tool_calls"),
structured_messages=_opaque_dict_list(inputs.get("structured_messages")),
tools=_opaque_dict_list(inputs.get("tools")),
tool_calls=_opaque_dict_list(inputs.get("tool_calls")),
)
if input_type == "request":
@ -305,9 +401,8 @@ class StraikerGuardrail(CustomGuardrail):
response_obj = request_data.get("response")
content.finish_reason = _response_finish_reason(response_obj)
original_messages = request_data.get("messages")
request_content = StraikerWebhookContent(
structured_messages=original_messages if isinstance(original_messages, list) else None,
structured_messages=_opaque_dict_list(_request_structured_messages(request_data)),
)
phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none"
event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase))

View file

@ -147,8 +147,10 @@ class UnifiedLLMGuardrails(CustomLogger):
litellm_logging_obj=data.get("litellm_logging_obj"),
)
# Add guardrail to applied guardrails header
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name)
if not guardrail_to_apply.records_own_guardrail_information:
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=guardrail_to_apply.guardrail_name
)
return data
async def async_moderation_hook(
@ -274,8 +276,10 @@ class UnifiedLLMGuardrails(CustomLogger):
if e.original_response is None:
e.original_response = response
raise
# Add guardrail to applied guardrails header
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name)
if not guardrail_to_apply.records_own_guardrail_information:
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=guardrail_to_apply.guardrail_name
)
return response

View file

@ -489,6 +489,9 @@ class InMemoryGuardrailHandler:
"skip_tool_message_in_guardrail",
getattr(litellm_params, "skip_tool_message_in_guardrail", None),
)
configured_run_in_parallel = getattr(litellm_params, "run_in_parallel", None)
if configured_run_in_parallel is not None:
custom_guardrail_callback.run_in_parallel = bool(configured_run_in_parallel)
parsed_guardrail = Guardrail(
guardrail_id=guardrail.get("guardrail_id"),

View file

@ -0,0 +1,493 @@
"""
SAML 2.0 SSO for the LiteLLM proxy admin UI.
Supports both SP-initiated and IdP-initiated login via the HTTP-POST binding,
using the OneLogin python3-saml toolkit for signature, audience and time
validation. The IdP is configured from its metadata (``SAML_IDP_METADATA_URL``
or inline ``SAML_IDP_METADATA_XML``); a successful login is mapped to a
``CustomOpenID`` and handed to the shared post-login path used by every other
SSO provider.
python3-saml pulls in the native ``xmlsec``/``libxml2`` libraries, so it is an
optional dependency. When it is not installed the SAML routes return a clear
error instead of breaking proxy startup.
"""
# python3-saml ships no type stubs, so the type checker sees every onelogin call
# as Unknown and the guarded optional import as possibly-unbound. Values crossing
# that boundary are cast() to concrete types at each use site; these directives
# silence only the unavoidable noise from the untyped dependency in this module.
# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false
# pyright: reportUnknownArgumentType=false, reportUnknownParameterType=false
# pyright: reportMissingTypeStubs=false, reportPossiblyUnboundVariable=false
# pyright: reportConstantRedefinition=false
import asyncio
import hashlib
import os
import secrets
import time
from typing import cast
from urllib.parse import parse_qsl
from fastapi import HTTPException, Request, status
from fastapi.responses import RedirectResponse
from pydantic import ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
from litellm.proxy.management_endpoints.types import CustomOpenID, get_litellm_user_role
from litellm.proxy.utils import get_custom_url
try:
from onelogin.saml2.auth import OneLogin_Saml2_Auth
from onelogin.saml2.idp_metadata_parser import OneLogin_Saml2_IdPMetadataParser
from onelogin.saml2.settings import OneLogin_Saml2_Settings
from onelogin.saml2.xml_utils import OneLogin_Saml2_XML
SAML_AVAILABLE = True
except ImportError:
SAML_AVAILABLE = False
SAML_LOGIN_ROUTE = "sso/saml/login"
SAML_CALLBACK_ROUTE = "sso/saml/callback"
SAML_METADATA_ROUTE = "sso/saml/metadata"
_SAML_AUTHN_STATE_COOKIE = "litellm_saml_authn"
_SAML_IDP_SETTINGS_CACHE_PREFIX = "saml_idp_settings"
_SAML_AUTHN_REQUEST_CACHE_PREFIX = "saml_authn_request"
_SAML_CONSUMED_ASSERTION_CACHE_PREFIX = "saml_consumed_assertion"
_SAML_AUTHN_REQUEST_TTL_SECONDS = 600
_SAML_IDP_METADATA_TTL_SECONDS = 3600
_SAML_METADATA_FETCH_TIMEOUT_SECONDS = 10
_SAML_MAX_POST_BYTES = 5 * 1024 * 1024
# The replay guard tracks each assertion's NotOnOrAfter so it spans the full
# validity window; the floor covers IdPs that issue hour-long assertions or omit
# the timestamp, and the cap bounds cache growth.
_SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS = 3600
_SAML_REPLAY_GUARD_MAX_TTL_SECONDS = 86400
_EMAIL_ATTRIBUTE_CANDIDATES = (
"urn:oid:0.9.2342.19200300.100.1.3",
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/emailaddress",
"email",
"emailAddress",
"mail",
"Email",
)
_FIRST_NAME_ATTRIBUTE_CANDIDATES = (
"urn:oid:2.5.4.42",
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/givenname",
"givenName",
"first_name",
"firstName",
)
_LAST_NAME_ATTRIBUTE_CANDIDATES = (
"urn:oid:2.5.4.4",
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/surname",
"sn",
"surname",
"last_name",
"lastName",
)
_ROLE_ATTRIBUTE_CANDIDATES = ("role", "roles", "litellm_role")
_TEAM_IDS_ATTRIBUTE_CANDIDATES = ("teams", "team_ids", "groups")
def _saml_unavailable_error() -> HTTPException:
return HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail=(
"SAML SSO requires the optional 'python3-saml' dependency, which is "
"not installed. Re-install litellm with the saml extra: "
"'pip install litellm[saml]'. The saml extra bundles the native "
"xmlsec/libxml2 libraries, so no system packages are required."
),
)
class SAMLAuthHandler:
"""SP- and IdP-initiated SAML 2.0 login for the admin UI."""
@staticmethod
def _env(name: str, default: str | None = None) -> str | None:
return os.getenv(name, default)
@staticmethod
def is_saml_configured() -> bool:
return bool(SAMLAuthHandler._env("SAML_IDP_METADATA_URL") or SAMLAuthHandler._env("SAML_IDP_METADATA_XML"))
@staticmethod
def _bool_env(name: str, default: bool) -> bool:
raw = SAMLAuthHandler._env(name)
if raw is None:
return default
return raw.strip().lower() in ("true", "1", "yes", "on")
@staticmethod
def _base_url(request: Request) -> str:
base = get_custom_url(request_base_url=str(request.base_url))
return base if base.endswith("/") else base + "/"
@staticmethod
def _is_https(request: Request) -> bool:
return SAMLAuthHandler._base_url(request).startswith("https")
@staticmethod
def _acs_url(request: Request) -> str:
return SAMLAuthHandler._base_url(request) + SAML_CALLBACK_ROUTE
@staticmethod
def _metadata_url(request: Request) -> str:
return SAMLAuthHandler._base_url(request) + SAML_METADATA_ROUTE
@staticmethod
def _sp_entity_id(request: Request) -> str:
return SAMLAuthHandler._env("SAML_SP_ENTITY_ID") or SAMLAuthHandler._metadata_url(request)
@staticmethod
async def _load_idp_settings(cache: DualCache) -> dict[str, object]:
metadata_url = SAMLAuthHandler._env("SAML_IDP_METADATA_URL")
metadata_xml = SAMLAuthHandler._env("SAML_IDP_METADATA_XML")
source = metadata_url or metadata_xml
if source is None:
raise HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail="SAML SSO is not configured. Set SAML_IDP_METADATA_URL or SAML_IDP_METADATA_XML.",
)
cache_key = f"{_SAML_IDP_SETTINGS_CACHE_PREFIX}:{hashlib.sha256(source.encode()).hexdigest()}"
cached = cache.get_cache(key=cache_key)
if isinstance(cached, dict):
return cast(dict[str, object], cached) # cast-ok: untyped python3-saml
if metadata_url is not None:
parsed = await asyncio.to_thread(
OneLogin_Saml2_IdPMetadataParser.parse_remote,
metadata_url,
validate_cert=SAMLAuthHandler._bool_env("SAML_IDP_METADATA_VALIDATE_CERT", True),
timeout=_SAML_METADATA_FETCH_TIMEOUT_SECONDS,
)
else:
parsed = OneLogin_Saml2_IdPMetadataParser.parse(cast(str, metadata_xml)) # cast-ok: untyped python3-saml
idp_settings = cast(dict[str, object], parsed) # cast-ok: untyped python3-saml
if not idp_settings.get("idp"):
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="Could not parse an IdP entityID/SSO URL/certificate from the SAML metadata.",
)
cache.set_cache(key=cache_key, value=idp_settings, ttl=_SAML_IDP_METADATA_TTL_SECONDS)
return idp_settings
@staticmethod
def _build_settings(request: Request, idp_settings: dict[str, object]) -> dict[str, object]:
sp_settings: dict[str, object] = {
"strict": SAMLAuthHandler._bool_env("SAML_STRICT", True),
"debug": False,
"sp": {
"entityId": SAMLAuthHandler._sp_entity_id(request),
"assertionConsumerService": {
"url": SAMLAuthHandler._acs_url(request),
"binding": "urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST",
},
"NameIDFormat": SAMLAuthHandler._env(
"SAML_SP_NAME_ID_FORMAT",
"urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress",
),
},
"security": {
"wantAssertionsSigned": SAMLAuthHandler._bool_env("SAML_WANT_ASSERTIONS_SIGNED", True),
"wantMessagesSigned": SAMLAuthHandler._bool_env("SAML_WANT_MESSAGES_SIGNED", False),
"authnRequestsSigned": SAMLAuthHandler._bool_env("SAML_AUTHN_REQUESTS_SIGNED", False),
"wantNameId": True,
"requestedAuthnContext": False,
"rejectUnsolicitedResponsesWithInResponseTo": False,
},
}
return OneLogin_Saml2_IdPMetadataParser.merge_settings(sp_settings, idp_settings)
@staticmethod
def _prepare_request_data(request: Request, post_data: dict[str, str] | None = None) -> dict[str, object]:
base = SAMLAuthHandler._base_url(request)
scheme, _, host_part = base.partition("://")
host = host_part.split("/", 1)[0]
return {
"https": "on" if scheme == "https" else "off",
"http_host": host,
"script_name": "/" + SAML_CALLBACK_ROUTE,
"get_data": dict(request.query_params),
"post_data": post_data or {},
}
@staticmethod
async def _build_auth(
request: Request,
cache: DualCache,
post_data: dict[str, str] | None = None,
) -> "OneLogin_Saml2_Auth":
if not SAML_AVAILABLE:
raise _saml_unavailable_error()
idp_settings = await SAMLAuthHandler._load_idp_settings(cache)
settings = SAMLAuthHandler._build_settings(request, idp_settings)
request_data = SAMLAuthHandler._prepare_request_data(request, post_data)
try:
return OneLogin_Saml2_Auth(request_data, old_settings=settings)
except Exception as e: # noqa: BLE001 - toolkit exposes no common exception base; fail closed
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Invalid SAML configuration: {e}",
)
@staticmethod
async def build_login_redirect(
request: Request, cache: DualCache, relay_state: str | None = None
) -> RedirectResponse:
auth = await SAMLAuthHandler._build_auth(request, cache)
redirect_url = cast(str, auth.login(return_to=relay_state)) # cast-ok: untyped python3-saml
response = RedirectResponse(url=redirect_url, status_code=303)
request_id = cast(str | None, auth.get_last_request_id()) # cast-ok: untyped python3-saml
if request_id is not None:
cache.set_cache(
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}",
value="1",
ttl=_SAML_AUTHN_REQUEST_TTL_SECONDS,
)
secure = SAMLAuthHandler._is_https(request)
response.set_cookie(
key=_SAML_AUTHN_STATE_COOKIE,
value=request_id,
max_age=_SAML_AUTHN_REQUEST_TTL_SECONDS,
httponly=True,
secure=secure,
samesite="none" if secure else "lax",
)
return response
@staticmethod
async def build_sp_metadata(request: Request, cache: DualCache) -> str:
if not SAML_AVAILABLE:
raise _saml_unavailable_error()
idp_settings = await SAMLAuthHandler._load_idp_settings(cache)
settings = SAMLAuthHandler._build_settings(request, idp_settings)
saml_settings = OneLogin_Saml2_Settings(settings, sp_validation_only=True)
metadata = cast(str, saml_settings.get_sp_metadata()) # cast-ok: untyped python3-saml
errors = cast(list[str], saml_settings.validate_metadata(metadata)) # cast-ok: untyped python3-saml
if errors:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail=f"Invalid SP metadata: {', '.join(errors)}",
)
return metadata
@staticmethod
async def read_acs_post_data(request: Request) -> dict[str, str]:
"""Read the ACS POST form under a hard size cap before any base64/XML decoding.
Bounds both Content-Length-declared and chunked requests so an unauthenticated
caller cannot force unbounded buffering while decoding the SAMLResponse."""
declared = request.headers.get("content-length")
if declared is not None and declared.isdigit() and int(declared) > _SAML_MAX_POST_BYTES:
raise HTTPException(
status_code=status.HTTP_413_CONTENT_TOO_LARGE,
detail="SAML response exceeds the maximum allowed size.",
)
body = bytearray()
async for chunk in request.stream():
body += chunk
if len(body) > _SAML_MAX_POST_BYTES:
raise HTTPException(
status_code=status.HTTP_413_CONTENT_TOO_LARGE,
detail="SAML response exceeds the maximum allowed size.",
)
return dict(parse_qsl(body.decode("utf-8", "replace")))
@staticmethod
async def handle_acs(request: Request, cache: DualCache, post_data: dict[str, str]) -> CustomOpenID:
auth = await SAMLAuthHandler._build_auth(request, cache, post_data=post_data)
browser_request_id = request.cookies.get(_SAML_AUTHN_STATE_COOKIE)
try:
auth.process_response(request_id=browser_request_id)
except Exception as e: # noqa: BLE001 - toolkit exposes no common exception base; fail closed
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=f"Could not process SAML response: {e}",
)
errors = cast(list[str], auth.get_errors()) # cast-ok: untyped python3-saml
if errors or not auth.is_authenticated():
reason = auth.get_last_error_reason()
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=f"SAML authentication failed: {reason or ', '.join(errors)}",
)
await SAMLAuthHandler._enforce_response_binding(auth, cache, browser_request_id)
return SAMLAuthHandler._result_from_auth(auth)
@staticmethod
def _replay_guard_ttl(auth: "OneLogin_Saml2_Auth") -> int:
not_on_or_after = auth.get_last_assertion_not_on_or_after()
if not isinstance(not_on_or_after, int):
return _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS
remaining = not_on_or_after - int(time.time())
return min(
max(remaining, _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS),
_SAML_REPLAY_GUARD_MAX_TTL_SECONDS,
)
@staticmethod
def _response_in_response_to(auth: "OneLogin_Saml2_Auth") -> str | None:
"""The request id this response answers, read from the Response element or, when the
IdP only stamps it on the bearer SubjectConfirmationData, from there. A non-None value
marks the response as solicited (SP-initiated) and so requiring browser binding."""
value = cast(str | None, auth.get_last_response_in_response_to()) # cast-ok: untyped python3-saml
if value:
return value
xml = cast(bytes | None, auth.get_last_response_xml()) # cast-ok: untyped python3-saml
if not xml:
return None
root = OneLogin_Saml2_XML.to_etree(xml)
for node in OneLogin_Saml2_XML.query(root, "//saml:SubjectConfirmationData[@InResponseTo]"):
irt = cast(str | None, node.get("InResponseTo")) # cast-ok: untyped python3-saml
if irt:
return irt
return None
@staticmethod
async def _enforce_response_binding(
auth: "OneLogin_Saml2_Auth",
cache: DualCache,
browser_request_id: str | None,
) -> None:
in_response_to = SAMLAuthHandler._response_in_response_to(auth)
if in_response_to is not None:
authn_key = f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{in_response_to}"
if cache.get_cache(key=authn_key) is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="SAML response references an unknown or already-used login request.",
)
if browser_request_id is None or not secrets.compare_digest(browser_request_id, in_response_to):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="SAML response is not bound to this browser's login request.",
)
elif browser_request_id is not None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="SAML response is not bound to this browser's login request.",
)
elif not SAMLAuthHandler._bool_env("SAML_ALLOW_UNSOLICITED", False):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Unsolicited (IdP-initiated) SAML responses are disabled.",
)
elif cache.redis_cache is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=(
"Unsolicited (IdP-initiated) SAML responses require a shared Redis cache "
"so the replay guard is enforced across every worker."
),
)
assertion_id = cast(str | None, auth.get_last_assertion_id()) # cast-ok: untyped python3-saml
if assertion_id is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="SAML assertion is missing the required ID attribute.",
)
consumed_key = f"{_SAML_CONSUMED_ASSERTION_CACHE_PREFIX}:{assertion_id}"
consumed_count = await cache.async_increment_cache(
key=consumed_key, value=1, ttl=SAMLAuthHandler._replay_guard_ttl(auth)
)
if consumed_count is not None and consumed_count > 1:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="SAML assertion has already been used (replay detected).",
)
@staticmethod
def _result_from_auth(auth: "OneLogin_Saml2_Auth") -> CustomOpenID:
attributes = cast(dict[str, list[str]], auth.get_attributes()) # cast-ok: untyped python3-saml
name_id = cast(str | None, auth.get_nameid()) # cast-ok: untyped python3-saml
email = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_EMAIL", _EMAIL_ATTRIBUTE_CANDIDATES)
if email is None and name_id is not None and "@" in name_id:
email = name_id
if email is None and SAMLAuthHandler._env("ALLOWED_EMAIL_DOMAINS") is not None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=(
"SAML assertion did not contain an email address, but ALLOWED_EMAIL_DOMAINS "
"restricts sign-in by email domain."
),
)
user_id = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_USER_ID", ()) or name_id or email
if user_id is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="SAML assertion did not contain a usable subject (NameID) or email.",
)
first_name = SAMLAuthHandler._attribute_value(
attributes, "SAML_ATTRIBUTE_FIRST_NAME", _FIRST_NAME_ATTRIBUTE_CANDIDATES
)
last_name = SAMLAuthHandler._attribute_value(
attributes, "SAML_ATTRIBUTE_LAST_NAME", _LAST_NAME_ATTRIBUTE_CANDIDATES
)
role_value = SAMLAuthHandler._attribute_value(attributes, "SAML_ATTRIBUTE_ROLE", _ROLE_ATTRIBUTE_CANDIDATES)
team_ids = SAMLAuthHandler._attribute_values(
attributes, "SAML_ATTRIBUTE_TEAM_IDS", _TEAM_IDS_ATTRIBUTE_CANDIDATES
)
display_name = " ".join(part for part in (first_name, last_name) if part) or email
verbose_proxy_logger.info(f"SAML login: subject={user_id}, email={email}, attributes={list(attributes.keys())}")
try:
return CustomOpenID(
id=user_id,
email=email,
first_name=first_name,
last_name=last_name,
display_name=display_name,
picture=None,
provider="saml",
team_ids=team_ids,
user_role=get_litellm_user_role(role_value) if role_value else None,
)
except ValidationError as e:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=f"SAML assertion contained an invalid subject or email: {e}",
)
@staticmethod
def _attribute_value(
attributes: dict[str, list[str]],
env_override: str,
candidates: tuple[str, ...],
) -> str | None:
values = SAMLAuthHandler._attribute_values(attributes, env_override, candidates)
return values[0] if values else None
@staticmethod
def _attribute_values(
attributes: dict[str, list[str]],
env_override: str,
candidates: tuple[str, ...],
) -> list[str]:
override = SAMLAuthHandler._env(env_override)
keys = (override, *candidates) if override else candidates
for key in keys:
values = attributes.get(key)
if values:
return [v for v in values if v]
return []

View file

@ -100,6 +100,7 @@ from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
from litellm.proxy.management_endpoints.sso.saml_sso import SAMLAuthHandler
from litellm.proxy.management_endpoints.sso_helper_utils import (
check_is_admin_only_access,
has_admin_ui_access,
@ -857,6 +858,27 @@ def process_sso_jwt_access_token(
return None
async def _raise_if_sso_exceeds_free_user_limit(premium_user: bool, prisma_client: PrismaClient | None) -> None:
"""Free tier allows SSO for up to 5 billable users; beyond that requires an Enterprise license."""
if premium_user is True:
return
if prisma_client is None:
raise ProxyException(
message=CommonProxyErrors.db_not_connected_error.value,
type=ProxyErrorTypes.auth_error,
param="premium_user",
code=status.HTTP_403_FORBIDDEN,
)
billable_users = await UserRepository(prisma_client).count_billable_users()
if billable_users and billable_users > 5:
raise ProxyException(
message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You configured SSO (one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, `GENERIC_CLIENT_ID`, or SAML) in your env. Please unset it",
type=ProxyErrorTypes.auth_error,
param="premium_user",
code=status.HTTP_403_FORBIDDEN,
)
@router.get("/sso/key/generate", tags=["experimental"], include_in_schema=False)
async def google_login(
request: Request,
@ -876,6 +898,7 @@ async def google_login(
general_settings,
premium_user,
prisma_client,
user_api_key_cache,
user_custom_ui_sso_sign_in_handler,
)
@ -891,25 +914,13 @@ async def google_login(
return admin_ui_disabled()
####### Check if user is a Enterprise / Premium User #######
if microsoft_client_id is not None or google_client_id is not None or generic_client_id is not None:
if premium_user is not True:
# Check if under 'free SSO user' limit
if prisma_client is not None:
billable_users = await UserRepository(prisma_client).count_billable_users()
if billable_users and billable_users > 5:
raise ProxyException(
message="You must be a LiteLLM Enterprise user to use SSO for more than 5 users. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this",
type=ProxyErrorTypes.auth_error,
param="premium_user",
code=status.HTTP_403_FORBIDDEN,
)
else:
raise ProxyException(
message=CommonProxyErrors.db_not_connected_error.value,
type=ProxyErrorTypes.auth_error,
param="premium_user",
code=status.HTTP_403_FORBIDDEN,
)
if (
microsoft_client_id is not None
or google_client_id is not None
or generic_client_id is not None
or SAMLAuthHandler.is_saml_configured()
):
await _raise_if_sso_exceeds_free_user_limit(premium_user, prisma_client)
####### Detect DB + MASTER KEY in .env #######
missing_env_vars = show_missing_vars_in_env()
@ -947,6 +958,19 @@ async def google_login(
"Enterprise features are not available. Custom UI SSO sign-in requires LiteLLM Enterprise."
)
if (
microsoft_client_id is None
and google_client_id is None
and generic_client_id is None
and SAMLAuthHandler.is_saml_configured()
):
verbose_proxy_logger.info("Redirecting to SAML SSO login")
return await SAMLAuthHandler.build_login_redirect(
request=request,
cache=user_api_key_cache,
relay_state=return_to,
)
# Check if we should use SSO handler
if (
SSOAuthenticationHandler.should_use_sso_handler(
@ -1913,6 +1937,81 @@ async def auth_callback(request: Request, state: Optional[str] = None):
)
@router.get("/sso/saml/login", tags=["experimental"], include_in_schema=False)
async def saml_login(request: Request, return_to: str | None = None):
"""SP-initiated SAML login. Redirects the user to the configured IdP."""
from litellm.proxy.proxy_server import user_api_key_cache
_disable_ui_flag = os.getenv("DISABLE_ADMIN_UI")
if _disable_ui_flag is not None and str_to_bool(value=_disable_ui_flag):
return admin_ui_disabled()
return await SAMLAuthHandler.build_login_redirect(request=request, cache=user_api_key_cache, relay_state=return_to)
@router.get("/sso/saml/metadata", tags=["experimental"], include_in_schema=False)
async def saml_metadata(request: Request):
"""Service Provider metadata XML, for registering this proxy at the IdP."""
from litellm.proxy.proxy_server import user_api_key_cache
metadata = await SAMLAuthHandler.build_sp_metadata(request=request, cache=user_api_key_cache)
return Response(content=metadata, media_type="application/xml")
@router.post("/sso/saml/callback", tags=["experimental"], include_in_schema=False)
async def saml_callback(request: Request):
"""Assertion Consumer Service. Validates the IdP assertion and issues a UI session."""
from litellm.proxy.proxy_server import (
general_settings,
jwt_handler,
master_key,
premium_user,
prisma_client,
user_api_key_cache,
)
_disable_ui_flag = os.getenv("DISABLE_ADMIN_UI")
if _disable_ui_flag is not None and str_to_bool(value=_disable_ui_flag):
return admin_ui_disabled()
if prisma_client is None:
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
if master_key is None:
raise ProxyException(
message="Master Key not set for Proxy. Set `LITELLM_MASTER_KEY` in .env or general_settings:master_key in config.yaml.",
type=ProxyErrorTypes.auth_error,
param="master_key",
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
)
post_data = await SAMLAuthHandler.read_acs_post_data(request)
if "SAMLResponse" not in post_data:
raise HTTPException(status_code=400, detail="Missing SAMLResponse in callback request.")
result = await SAMLAuthHandler.handle_acs(request=request, cache=user_api_key_cache, post_data=post_data)
await _raise_if_sso_exceeds_free_user_limit(premium_user, prisma_client)
ui_access_mode = general_settings.get("ui_access_mode", None)
relay_state = post_data.get("RelayState")
cp_return_to: str | None = (
relay_state
if isinstance(relay_state, str) and SSOAuthenticationHandler._validate_return_to(relay_state)
else None
)
return await SSOAuthenticationHandler.get_redirect_response_from_openid(
result=result,
request=request,
received_response=None,
generic_client_id=None,
ui_access_mode=ui_access_mode,
access_token_payload=None,
jwt_handler=jwt_handler,
return_to=cp_return_to,
)
async def _build_cli_sso_user_defined_values(
result: Union[OpenID, dict],
parsed_openid_result: ParsedOpenIDResult,

View file

@ -38,6 +38,10 @@ from litellm._uuid import uuid
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -668,16 +672,13 @@ def _carry_guardrail_logging_info(request_data: dict, guardrail_data: Optional[d
"""
if guardrail_data is None:
return
source_metadata = guardrail_data.get("metadata")
if not isinstance(source_metadata, dict):
return
source_key = get_metadata_variable_name_from_kwargs(guardrail_data)
source_metadata = guardrail_data.get(source_key) or {}
entries = source_metadata.get("standard_logging_guardrail_information")
if not entries:
return
metadata = request_data.get("metadata")
if not isinstance(metadata, dict):
metadata = request_data["metadata"] = {}
_, metadata = get_or_create_metadata_bucket(request_data)
metadata.setdefault("standard_logging_guardrail_information", list(entries))

View file

@ -1,5 +1,5 @@
import asyncio
from typing import TYPE_CHECKING, Any, Literal, Optional
from typing import TYPE_CHECKING, Any, Literal, Mapping, Optional
import httpx
from fastapi import HTTPException, status
@ -145,6 +145,30 @@ class ProxyModelNotFoundError(HTTPException):
super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail)
REQUIRED_BODY_PARAM_BY_ROUTE: Mapping[str, str] = {
"acompletion": "messages",
"aembedding": "input",
}
class ProxyMissingRequiredParamError(HTTPException):
def __init__(self, route: str, param: str):
detail = {"error": f"{route}: Missing required parameter: '{param}'."}
super().__init__(status_code=status.HTTP_400_BAD_REQUEST, detail=detail)
self.type = "invalid_request_error"
self.param = param
def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None:
required_param = REQUIRED_BODY_PARAM_BY_ROUTE.get(route_type)
if required_param is None or data.get(required_param) is not None:
return
raise ProxyMissingRequiredParamError(
route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type),
param=required_param,
)
def get_team_id_from_data(data: dict) -> Optional[str]:
"""
Get the team id from the data's metadata or litellm_metadata params.
@ -353,6 +377,8 @@ async def route_request(
"""
Common helper to route the request
"""
raise_if_required_body_param_missing(route_type=route_type, data=data)
await add_shared_session_to_data(data)
# Strip router-internal mock_testing_* flags. Combined with an

View file

@ -1400,6 +1400,14 @@ class ProxyLogging:
self._process_guardrail_metadata(data)
return data
parallel_guardrails: tuple[CustomGuardrail, ...] = tuple(
cb
for cb in caps.resolved_callbacks
if isinstance(cb, CustomGuardrail)
and getattr(cb, "run_in_parallel", False)
and not (cb.guardrail_name and cb.guardrail_name in pipeline_managed)
)
deferred_route_exc: Optional[SensitiveDataRouteException] = None
for _callback in caps.resolved_callbacks:
start_time = time.time()
@ -1409,6 +1417,9 @@ class ProxyLogging:
if _callback.guardrail_name and _callback.guardrail_name in pipeline_managed:
continue
if getattr(_callback, "run_in_parallel", False):
continue
result = await self._process_guardrail_callback(
callback=_callback,
data=data, # type: ignore
@ -1465,6 +1476,14 @@ class ProxyLogging:
if deferred_route_exc is not None and data is not None:
data = await self._handle_sensitive_data_route_exception(deferred_route_exc, data, user_api_key_dict)
if parallel_guardrails and data is not None:
await self._run_parallel_pre_call_guardrails(
guardrails=parallel_guardrails,
data=data,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
)
if data is not None:
self._process_guardrail_metadata(data)
@ -1477,6 +1496,47 @@ class ProxyLogging:
except Exception as e:
raise e
async def _run_parallel_pre_call_guardrails(
self,
guardrails: tuple[CustomGuardrail, ...],
data: dict,
user_api_key_dict: UserAPIKeyAuth,
call_type: CallTypesLiteral,
) -> None:
"""
Run opted-in pre_call guardrails concurrently against one shared payload
snapshot. These guardrails are declared block-only, so any modified data
they return is discarded; they run for their blocking side effect (raising
to reject the request before it reaches the LLM). Every guardrail is
awaited to completion (``return_exceptions=True``) so a raise by one never
leaves the others running as unobserved background tasks. A guardrail that
blocks (any exception other than a reroute or passthrough) takes precedence
over one that only changes the request flow, so a fast reroute can never
let a slower block be bypassed; the request is rejected before it reaches
the LLM, preserving the pre-call barrier that ``during_call`` guardrails
cannot provide. Per-guardrail latency is recorded by
``_process_guardrail_callback``'s own metrics.
"""
results = await asyncio.gather(
*(
self._process_guardrail_callback(
callback=callback,
data=data,
user_api_key_dict=user_api_key_dict,
call_type=call_type,
event_type=GuardrailEventHooks.pre_call,
)
for callback in guardrails
),
return_exceptions=True,
)
raised = tuple(result for result in results if isinstance(result, BaseException))
blocking = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None)
if blocking is not None:
raise blocking
if raised:
raise raised[0]
async def _handle_sensitive_data_route_exception(
self,
exc: SensitiveDataRouteException,
@ -2277,9 +2337,16 @@ class ProxyLogging:
# Merge model-level guardrails before checking which guardrails to run
guardrail_data = _check_and_merge_model_level_guardrails(data=data, llm_router=llm_router)
parallel_guardrails: tuple[CustomGuardrail, ...] = tuple(
callback for callback in guardrail_callbacks if getattr(callback, "run_in_parallel", False)
)
for callback in guardrail_callbacks:
# Main - V2 Guardrails implementation
if getattr(callback, "run_in_parallel", False):
continue
if (
callback.should_run_guardrail(
data=guardrail_data,
@ -2316,6 +2383,15 @@ class ProxyLogging:
if guardrail_response is not None:
response = guardrail_response
if parallel_guardrails:
await self._run_parallel_post_call_guardrails(
guardrails=parallel_guardrails,
data=data,
guardrail_data=guardrail_data,
response=response,
user_api_key_dict=user_api_key_dict,
)
############ Handle CustomLogger ###############################
#################################################################
@ -2329,6 +2405,65 @@ class ProxyLogging:
raise e
return response
async def _run_parallel_post_call_guardrails(
self,
guardrails: tuple[CustomGuardrail, ...],
data: dict,
guardrail_data: dict,
response: LLMResponseTypes,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""
Run opted-in post_call guardrails concurrently against the response
produced by the sequential guardrails. These guardrails are declared
block-only, so any modified response they return is discarded; they run
for their blocking side effect (raising to reject the response before it
reaches the client). Every guardrail is awaited to completion
(``return_exceptions=True``) so a raise by one never leaves the others
running as unobserved background tasks. A guardrail that blocks (any
exception other than a passthrough) takes precedence over one that only
changes the response flow, so a fast passthrough can never let a slower
block be bypassed. Each per-guardrail coroutine sets ``guardrail_to_apply``
immediately before awaiting, and the unified hook pops it before its first
suspension point, so concurrent guardrails never race on that key.
"""
async def _run_one(callback: CustomGuardrail) -> None:
if callback.should_run_guardrail(data=guardrail_data, event_type=GuardrailEventHooks.post_call) is not True:
return
if "apply_guardrail" in type(callback).__dict__:
data["guardrail_to_apply"] = callback
await self._run_guardrail_with_metrics(
callback,
unified_guardrail.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
),
"post_call",
)
else:
await self._run_guardrail_with_metrics(
callback,
callback.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
),
"post_call",
)
results = await asyncio.gather(
*(_run_one(callback) for callback in guardrails),
return_exceptions=True,
)
raised = tuple(result for result in results if isinstance(result, BaseException))
blocking = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None)
if blocking is not None:
raise blocking
if raised:
raise raised[0]
async def post_call_response_headers_hook(
self,
data: dict,

View file

@ -5,6 +5,7 @@ from typing import (
Dict,
Iterable,
List,
Mapping,
Optional,
Type,
Union,
@ -24,6 +25,7 @@ from litellm.types.llms.openai import (
ResponseInputParam,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamOptions,
ResponseText,
)
from litellm.types.responses.main import DecodedResponseId
@ -35,6 +37,17 @@ from litellm.types.utils import (
)
def normalize_responses_api_stream_options(
stream_options: object,
) -> ResponsesAPIStreamOptions | None:
if not isinstance(stream_options, Mapping):
return None
include_obfuscation = stream_options.get("include_obfuscation")
if not isinstance(include_obfuscation, bool):
return None
return ResponsesAPIStreamOptions(include_obfuscation=include_obfuscation)
class ResponsesAPIRequestUtils:
"""Helper utils for constructing ResponseAPI requests"""
@ -156,15 +169,19 @@ class ResponsesAPIRequestUtils:
drop_params=should_drop_params,
)
stream_options = normalize_responses_api_stream_options(mapped_params.get("stream_options"))
params_with_normalized_stream_options = {
**{key: value for key, value in mapped_params.items() if key != "stream_options"},
**({} if stream_options is None else {"stream_options": stream_options}),
}
# add any allowed_openai_params to the mapped_params
mapped_params = _apply_openai_param_overrides(
optional_params=mapped_params,
return _apply_openai_param_overrides(
optional_params=params_with_normalized_stream_options,
non_default_params=non_default_params,
allowed_openai_params=allowed_openai_params or [],
)
return mapped_params
@staticmethod
def get_requested_response_api_optional_param(
params: Dict[str, Any],

View file

@ -910,6 +910,17 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
),
)
run_in_parallel: Optional[bool] = Field(
default=None,
description=(
"When True, this pre_call or post_call guardrail runs concurrently with other opted-in "
"guardrails of the same hook, after the sequential guardrails have run. Use only for "
"block-only guardrails that inspect and reject; do not enable it for guardrails that "
"modify the request or response (e.g. PII masking or sensitive-data routing), since "
"parallel runs share one snapshot and their mutations would race."
),
)
@field_validator(
"mode",
"default_action",

View file

@ -1145,6 +1145,10 @@ class ContextManagementEntry(TypedDict, total=False):
"""Token threshold at which compaction is triggered for this entry. Minimum 1000."""
class ResponsesAPIStreamOptions(TypedDict, total=False):
include_obfuscation: bool
class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
"""TypedDict for Optional parameters supported by the responses API."""
@ -1171,7 +1175,7 @@ class ResponsesAPIOptionalRequestParams(TypedDict, total=False):
max_tool_calls: Optional[int]
prompt_cache_key: Optional[str]
prompt_cache_retention: Optional[str]
stream_options: Optional[dict]
stream_options: Optional[ResponsesAPIStreamOptions]
top_logprobs: Optional[int]
partial_images: Optional[int] # Number of partial images to generate (1-3) for streaming image generation
context_management: Optional[List[ContextManagementEntry]]

View file

@ -4,9 +4,6 @@ from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
from litellm.types.utils import ChatCompletionMessageToolCall
from .base import GuardrailConfigModel
StraikerWebhookEventType = Literal["pre_call", "post_call"]
@ -32,9 +29,9 @@ class StraikerWebhookContent(BaseModel):
texts: list[str] = Field(default_factory=list)
images: list[str] = Field(default_factory=list)
structured_messages: list[AllMessageValues] | None = None
structured_messages: list[dict[str, object]] | None = None
tools: list[dict[str, object]] | None = None
tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None
tool_calls: list[dict[str, object]] | None = None
finish_reason: str | None = None
@ -45,6 +42,7 @@ class StraikerWebhookUsage(BaseModel):
class StraikerWebhookContext(BaseModel):
call_surface: str
mode: list[str] | None = None
model: str | None = None
model_provider: str | None = None
destination: str | None = None

View file

@ -153,6 +153,24 @@ class SSOConfig(LiteLLMPydanticObjectBase):
description="Space-separated OAuth scopes requested from the generic provider, e.g. 'openid email profile'",
)
# SAML SSO
saml_idp_metadata_url: Optional[str] = Field(
default=None,
description="URL of the SAML IdP metadata to fetch and parse for SSO authentication",
)
saml_idp_metadata_xml: Optional[str] = Field(
default=None,
description="Inline SAML IdP metadata XML, used when a metadata URL is not available",
)
saml_sp_entity_id: Optional[str] = Field(
default=None,
description="SAML Service Provider entityID; defaults to the proxy's /sso/saml/metadata URL",
)
saml_allow_unsolicited: Optional[str] = Field(
default=None,
description="'true' to accept IdP-initiated (unsolicited) SAML responses, which cannot be browser-bound against login CSRF",
)
# Common settings
proxy_base_url: Optional[str] = Field(
default=None,

View file

@ -99,6 +99,10 @@ utils = [
"numpydoc>=1.8.0,<2.0",
]
caching = ["diskcache>=5.6.3,<6.0"]
# SAML SSO for the admin UI. python3-saml pulls in xmlsec/lxml, whose wheels
# bundle the native libxmlsec1/libxml2 libraries, so no system packages are
# required. Kept out of the base `proxy` extra so it stays optional.
saml = ["python3-saml>=1.16.0,<2.0"]
semantic-router = [
"semantic-router>=0.1.15,<1.0; python_version < '3.14'",
"aurelio-sdk>=0.0.19,<1.0; python_version < '3.14'",

View file

@ -222,7 +222,7 @@
"limit": 38
},
"RET504": {
"limit": 721
"limit": 719
},
"RUF010": {
"limit": 874
@ -324,7 +324,7 @@
"limit": 883
},
"UP006": {
"limit": 12792
"limit": 12789
},
"UP007": {
"limit": 2570

View file

@ -7,6 +7,14 @@
assertions: [succeeds]
source: "server.py:637"
rationale: Core operation; most common auth path; high usage
- id: mcp.list_tools.api_key.access_group_scoped
module: mcp
tier: P1
operation: list_tools
auth_family: api_key
assertions: [access_group_scoped]
source: "test_mcp_access_group_e2e.py"
rationale: "A key granted an MCP access group sees the tagged server's tools; a key with a different group does not. Access-group-scoped tool selection at key creation"
- id: mcp.list_tools.api_key.denied_without_permission
module: mcp
tier: P0

View file

@ -30,7 +30,12 @@ def assert_dd_mcp_creds() -> None:
)
def register_datadog_mcp(client: McpClient, resources: ResourceManager) -> str:
def register_datadog_mcp(
client: McpClient,
resources: ResourceManager,
*,
mcp_access_groups: list[str] | None = None,
) -> str:
assert_dd_mcp_creds()
name = f"e2e_dd_mcp_{unique_marker()}"
server_id = client.register_server(
@ -43,6 +48,7 @@ def register_datadog_mcp(client: McpClient, resources: ResourceManager) -> str:
"DD-APPLICATION-KEY": _dd_app_key(),
},
allowed_tools=[SEARCH_LOGS_TOOL],
mcp_access_groups=mcp_access_groups,
)
resources.defer(lambda: client.delete_server(server_id))
return server_id

View file

@ -36,6 +36,7 @@ class McpServerNewBody(BaseModel):
auth_type: str | None = None
static_headers: dict[str, str] | None = None
allowed_tools: list[str] | None = None
mcp_access_groups: list[str] | None = None
class McpServerNewResponse(BaseModel):
@ -155,6 +156,7 @@ class McpClient:
auth_type: str | None = None,
static_headers: dict[str, str] | None = None,
allowed_tools: list[str] | None = None,
mcp_access_groups: list[str] | None = None,
) -> str:
return unwrap(
self.proxy.transport.post(
@ -168,6 +170,7 @@ class McpClient:
auth_type=auth_type,
static_headers=static_headers,
allowed_tools=allowed_tools,
mcp_access_groups=mcp_access_groups,
),
response_type=McpServerNewResponse,
)
@ -196,10 +199,13 @@ class McpClient:
*,
user_id: str,
mcp_servers: list[str] | None,
mcp_access_groups: list[str] | None = None,
models: list[str] | None = None,
) -> str:
object_permission = (
ObjectPermission(mcp_servers=mcp_servers) if mcp_servers is not None else None
ObjectPermission(mcp_servers=mcp_servers, mcp_access_groups=mcp_access_groups)
if mcp_servers is not None or mcp_access_groups is not None
else None
)
return self.proxy.generate_key(
KeyGenerateBody(

View file

@ -0,0 +1,58 @@
"""Live e2e: MCP tool selection via access group at key creation.
An admin registers the Datadog remote MCP server tagged with a server-side
access group (`mcp_access_groups`). A key minted with that access group
(`object_permission.mcp_access_groups`) sees the server's tools; a key minted
with a different group does not. This exercises access-group-scoped tool
selection, the enterprise MCP surface where keys are granted tool access groups
rather than explicit server ids.
A tools/list that leaks the server across the access-group boundary fails hard.
Requires DD_API_KEY + DD_APP_KEY (the suite's real MCP upstream).
"""
import pytest
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
from e2e_config import unique_marker
from e2e_http import unwrap
from lifecycle import ResourceManager
from mcp_client import McpClient
pytestmark = pytest.mark.e2e
class TestMcpAccessGroupToolSelection:
@pytest.mark.covers("mcp.list_tools.api_key.access_group_scoped")
def test_access_group_scopes_tool_selection(
self, client: McpClient, resources: ResourceManager
) -> None:
group = f"e2e-mcp-grp-{unique_marker()}"
server_id = register_datadog_mcp(client, resources, mcp_access_groups=[group])
granted = client.generate_key(
user_id=f"e2e-mcp-ag-granted-{unique_marker()}",
mcp_servers=None,
mcp_access_groups=[group],
)
resources.defer(lambda: client.proxy.delete_key(granted))
other = client.generate_key(
user_id=f"e2e-mcp-ag-other-{unique_marker()}",
mcp_servers=None,
mcp_access_groups=[f"e2e-mcp-grp-absent-{unique_marker()}"],
)
resources.defer(lambda: client.proxy.delete_key(other))
granted_tools = unwrap(client.list_tools(granted))
assert granted_tools.tool_name_containing(server_id, SEARCH_LOGS_TOOL) is not None, (
f"key granted access group {group} did not see the tagged server's tool "
f"(upstream dead or access-group grant not applied): "
f"{granted_tools.tool_names_for_server(server_id)}"
)
other_tools = unwrap(client.list_tools(other)).tool_names_for_server(server_id)
assert other_tools == frozenset(), (
f"key with a different access group saw the server's tools; access-group tool "
f"selection leaked across the boundary: {other_tools}"
)

View file

@ -47,6 +47,7 @@ class KeyMetadata(BaseModel):
class ObjectPermission(BaseModel):
mcp_servers: list[str] | None = None
mcp_access_groups: list[str] | None = None
class KeyGenerateBody(BaseModel):

View file

@ -7,10 +7,12 @@ from typing import Any, Dict, List, Optional, Union
from unittest.mock import Mock
import pytest
from fastapi import Request
from fastapi import HTTPException, Request
from starlette.datastructures import State
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.utils import _get_docs_url, _get_openapi_url, _get_redoc_url
from litellm.types.guardrails import GuardrailEventHooks
sys.path.insert(
0, os.path.abspath("../..")
@ -2638,6 +2640,463 @@ async def test_during_call_hook_parallel_execution_with_error():
litellm.callbacks = original_callbacks
class _PreCallGuardrail(CustomGuardrail):
"""Test double for pre_call guardrails; records timing and observed payload."""
def __init__(self, name, run_in_parallel, execution_order, sleep=0.1, default_on=True):
super().__init__(
guardrail_name=name,
event_hook=GuardrailEventHooks.pre_call,
default_on=default_on,
run_in_parallel=run_in_parallel,
)
self.name = name
self.sleep = sleep
self.execution_order = execution_order
self.observed_content = None
self.was_called = False
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
self.was_called = True
self.observed_content = data["messages"][0]["content"]
self.execution_order.append(f"{self.name}_start")
await asyncio.sleep(self.sleep)
self.execution_order.append(f"{self.name}_end")
return None
@pytest.mark.asyncio
async def test_pre_call_hook_runs_opted_in_guardrails_in_parallel():
"""run_in_parallel pre_call guardrails execute concurrently (all start before any ends)."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
execution_order = []
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
try:
litellm.callbacks = [
_PreCallGuardrail(f"g{i}", run_in_parallel=True, execution_order=execution_order) for i in range(3)
]
result = await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
call_type="completion",
)
first_end_idx = next(i for i, item in enumerate(execution_order) if "end" in item)
starts_before_first_end = sum(1 for item in execution_order[:first_end_idx] if "start" in item)
assert starts_before_first_end == 3, f"expected 3 concurrent starts, got {starts_before_first_end}"
assert result["model"] == "gpt-4"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_pre_call_hook_runs_default_guardrails_sequentially():
"""Guardrails without run_in_parallel keep the sequential, one-at-a-time behavior."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
execution_order = []
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
try:
litellm.callbacks = [
_PreCallGuardrail(f"g{i}", run_in_parallel=False, execution_order=execution_order) for i in range(2)
]
start = asyncio.get_event_loop().time()
await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
call_type="completion",
)
elapsed = asyncio.get_event_loop().time() - start
assert execution_order == ["g0_start", "g0_end", "g1_start", "g1_end"]
assert elapsed >= 0.18, f"sequential run took {elapsed}s, expected >= 0.18s"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_pre_call_hook_sequential_mutations_precede_parallel_batch():
"""Sequential (mutating) guardrails run before the parallel batch, which sees their changes."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
execution_order = []
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
class MaskingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="masker",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=False,
)
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
data["messages"][0]["content"] = "MASKED"
return data
parallel_observer = _PreCallGuardrail("observer", run_in_parallel=True, execution_order=execution_order)
try:
litellm.callbacks = [parallel_observer, MaskingGuardrail()]
await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "secret"}]},
call_type="completion",
)
assert parallel_observer.observed_content == "MASKED"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_pre_call_hook_parallel_guardrail_blocks_request():
"""A raising parallel guardrail blocks the request before it reaches the LLM."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
class BlockingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="blocker",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=True,
)
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
raise HTTPException(status_code=400, detail="blocked by guardrail")
try:
litellm.callbacks = [BlockingGuardrail()]
with pytest.raises(HTTPException) as exc_info:
await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
call_type="completion",
)
assert exc_info.value.status_code == 400
assert "blocked by guardrail" in str(exc_info.value.detail)
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_pre_call_hook_parallel_guardrail_skipped_when_should_not_run():
"""A parallel guardrail that should_run_guardrail rejects is never invoked."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
execution_order = []
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
try:
guardrail = _PreCallGuardrail(
"off_by_default", run_in_parallel=True, execution_order=execution_order, default_on=False
)
litellm.callbacks = [guardrail]
result = await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
call_type="completion",
)
assert guardrail.was_called is False
assert result["model"] == "gpt-4"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_pre_call_hook_parallel_block_wins_over_reroute():
"""A slower block must win over a faster reroute so crafted input cannot bypass a block."""
from litellm.caching.caching import DualCache
from litellm.exceptions import SensitiveDataRouteException
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
class FastRerouteGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="rerouter",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=True,
)
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
raise SensitiveDataRouteException(route_to_model="on-prem", session_id="s1", guardrail_name="rerouter")
class SlowBlockingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="blocker",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=True,
)
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
await asyncio.sleep(0.1)
raise HTTPException(status_code=400, detail="blocked by guardrail")
try:
litellm.callbacks = [FastRerouteGuardrail(), SlowBlockingGuardrail()]
with pytest.raises(HTTPException) as exc_info:
await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
call_type="completion",
)
assert exc_info.value.status_code == 400
assert "blocked by guardrail" in str(exc_info.value.detail)
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_pre_call_hook_parallel_awaits_all_when_one_blocks():
"""A block must not orphan sibling guardrails; every parallel guardrail runs to completion."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
completed = []
class FastBlockingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="fast_blocker",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=True,
)
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
raise HTTPException(status_code=400, detail="blocked")
class SlowGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="slow",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=True,
)
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
await asyncio.sleep(0.1)
completed.append("slow")
return None
try:
litellm.callbacks = [FastBlockingGuardrail(), SlowGuardrail()]
with pytest.raises(HTTPException):
await proxy_logging.pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
call_type="completion",
)
assert completed == ["slow"], "slow guardrail was orphaned instead of awaited to completion"
finally:
litellm.callbacks = original_callbacks
class _PostCallGuardrail(CustomGuardrail):
"""Test double for post_call guardrails; records timing and invocation."""
def __init__(self, name, run_in_parallel, execution_order, sleep=0.1, default_on=True):
super().__init__(
guardrail_name=name,
event_hook=GuardrailEventHooks.post_call,
default_on=default_on,
run_in_parallel=run_in_parallel,
)
self.name = name
self.sleep = sleep
self.execution_order = execution_order
self.was_called = False
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
self.was_called = True
self.execution_order.append(f"{self.name}_start")
await asyncio.sleep(self.sleep)
self.execution_order.append(f"{self.name}_end")
return None
@pytest.mark.asyncio
async def test_post_call_hook_runs_opted_in_guardrails_in_parallel():
"""run_in_parallel post_call guardrails execute concurrently (all start before any ends)."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
execution_order = []
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
try:
litellm.callbacks = [
_PostCallGuardrail(f"g{i}", run_in_parallel=True, execution_order=execution_order) for i in range(3)
]
await proxy_logging.post_call_success_hook(
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
response=litellm.ModelResponse(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
)
first_end_idx = next(i for i, item in enumerate(execution_order) if "end" in item)
starts_before_first_end = sum(1 for item in execution_order[:first_end_idx] if "start" in item)
assert starts_before_first_end == 3, f"expected 3 concurrent starts, got {starts_before_first_end}"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_post_call_hook_runs_default_guardrails_sequentially():
"""post_call guardrails without run_in_parallel keep the sequential behavior."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
execution_order = []
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
try:
litellm.callbacks = [
_PostCallGuardrail(f"g{i}", run_in_parallel=False, execution_order=execution_order) for i in range(2)
]
start = asyncio.get_event_loop().time()
await proxy_logging.post_call_success_hook(
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
response=litellm.ModelResponse(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
)
elapsed = asyncio.get_event_loop().time() - start
assert execution_order == ["g0_start", "g0_end", "g1_start", "g1_end"]
assert elapsed >= 0.18, f"sequential run took {elapsed}s, expected >= 0.18s"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_post_call_hook_parallel_guardrail_blocks_response():
"""A raising parallel post_call guardrail blocks the response before it reaches the client."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
class BlockingPostCallGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="post_blocker",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
run_in_parallel=True,
)
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
raise HTTPException(status_code=400, detail="blocked response by guardrail")
try:
litellm.callbacks = [BlockingPostCallGuardrail()]
with pytest.raises(HTTPException) as exc_info:
await proxy_logging.post_call_success_hook(
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
response=litellm.ModelResponse(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
)
assert exc_info.value.status_code == 400
assert "blocked response by guardrail" in str(exc_info.value.detail)
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_post_call_hook_parallel_awaits_all_when_one_blocks():
"""A blocking post_call guardrail must not orphan its siblings; all run to completion."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
completed = []
class FastBlockingPostCall(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="fast_post_blocker",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
run_in_parallel=True,
)
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
raise HTTPException(status_code=400, detail="blocked")
class SlowPostCall(CustomGuardrail):
def __init__(self):
super().__init__(
guardrail_name="slow_post",
event_hook=GuardrailEventHooks.post_call,
default_on=True,
run_in_parallel=True,
)
async def async_post_call_success_hook(self, data, user_api_key_dict, response):
await asyncio.sleep(0.1)
completed.append("slow")
return None
try:
litellm.callbacks = [FastBlockingPostCall(), SlowPostCall()]
with pytest.raises(HTTPException):
await proxy_logging.post_call_success_hook(
data={"model": "gpt-4", "messages": [{"role": "user", "content": "hi"}]},
response=litellm.ModelResponse(),
user_api_key_dict=UserAPIKeyAuth(api_key="test_key", user_id="test_user"),
)
assert completed == ["slow"], "slow post_call guardrail was orphaned instead of awaited to completion"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
async def test_handle_logging_proxy_only_error_preserves_pass_through_call_type():
"""Ensure _handle_logging_proxy_only_error does not overwrite call_type

View file

@ -22,78 +22,6 @@ def redis_no_ping():
yield
@pytest.mark.parametrize("namespace", [None, "test"])
@pytest.mark.asyncio
async def test_redis_cache_async_increment(namespace, monkeypatch, redis_no_ping):
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache(namespace=namespace)
# Create an AsyncMock for the Redis client
mock_redis_instance = AsyncMock()
# Make sure the mock can be used as an async context manager
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
assert redis_cache is not None
expected_key = "test:test" if namespace else "test"
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
# Call async_set_cache
await redis_cache.async_increment(key=expected_key, value=1)
# Verify that the set method was called on the mock Redis instance
mock_redis_instance.incrbyfloat.assert_called_once_with(
name=expected_key, amount=1
)
@pytest.mark.asyncio
async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl(
monkeypatch, redis_no_ping
):
"""With refresh_ttl=True, every increment should call expire() to bump
the TTL, even when the key already has a TTL (counter-style use)."""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
mock_redis_instance.ttl.return_value = 42 # key already has ~42s left
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
await redis_cache.async_increment(
key="spend:team_member:u:t", value=0.05, refresh_ttl=True
)
mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60)
@pytest.mark.asyncio
async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl(
monkeypatch, redis_no_ping
):
"""Default (refresh_ttl=False) preserves window-style semantics: TTL is
set only on first creation, never refreshed (used by rate-limit windows)."""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
mock_redis_instance.ttl.return_value = 42 # key already has ~42s left
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
await redis_cache.async_increment(key="rate_limit:window", value=1)
mock_redis_instance.expire.assert_not_awaited()
@pytest.mark.parametrize("namespace", [None, "litellm"])
@pytest.mark.asyncio
async def test_async_delete_cache_applies_namespace(
@ -140,42 +68,6 @@ async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping)
assert client.connection_pool.connection_kwargs["socket_timeout"] == 1.0
@pytest.mark.asyncio
async def test_redis_cache_async_batch_get_cache(monkeypatch, redis_no_ping):
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
# Create an AsyncMock for the Redis client
mock_redis_instance = AsyncMock()
# Make sure the mock can be used as an async context manager
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
# Setup the return value for mget
mock_redis_instance.mget.return_value = [
b'{"key1": "value1"}',
None,
b'{"key3": "value3"}',
]
test_keys = ["key1", "key2", "key3"]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
# Call async_batch_get_cache
result = await redis_cache.async_batch_get_cache(key_list=test_keys)
# Verify mget was called with the correct keys
mock_redis_instance.mget.assert_called_once()
# Check that results were properly decoded
assert result["key1"] == {"key1": "value1"}
assert result["key2"] is None
assert result["key3"] == {"key3": "value3"}
@pytest.mark.asyncio
async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
"""Test the helper method that handles LPOP with count for Redis versions < 7.0"""
@ -202,41 +94,6 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
assert mock_pipeline.execute.call_count == 2
@pytest.mark.asyncio
async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping):
"""Verify that multiple rpush ops are batched into a single pipeline execute"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.rpush = MagicMock()
mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineRpushOperation
rpush_list = [
RedisPipelineRpushOperation(key="key1", values=["a", "b"]),
RedisPipelineRpushOperation(key="key2", values=["c"]),
RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]),
]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
assert result == [3, 5, 1]
assert mock_pipeline.rpush.call_count == 3
mock_pipeline.rpush.assert_any_call("key1", "a", "b")
mock_pipeline.rpush.assert_any_call("key2", "c")
mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f")
mock_pipeline.execute.assert_called_once()
@pytest.mark.asyncio
async def test_async_rpush_pipeline_empty_list_returns_empty(
monkeypatch, redis_no_ping
@ -256,183 +113,6 @@ async def test_async_rpush_pipeline_empty_list_returns_empty(
mock_redis_instance.pipeline.assert_not_called()
@pytest.mark.asyncio
async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping):
"""Pipeline errors should propagate"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.rpush = MagicMock()
mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down"))
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineRpushOperation
rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
with pytest.raises(ConnectionError, match="Redis down"):
await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
@pytest.mark.asyncio
async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping):
"""Verify that multiple lpop ops are batched into a single pipeline execute"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "7.0.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
mock_pipeline.execute = AsyncMock(
return_value=[
[b"val1", b"val2"], # key1 results
None, # key2 empty
[b"val3"], # key3 results
]
)
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [
RedisPipelineLpopOperation(key="key1", count=10),
RedisPipelineLpopOperation(key="key2", count=10),
RedisPipelineLpopOperation(key="key3", count=5),
]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
assert len(results) == 3
assert results[0] == ["val1", "val2"]
assert results[1] is None
assert results[2] == ["val3"]
mock_pipeline.execute.assert_called_once()
@pytest.mark.asyncio
async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results(
monkeypatch, redis_no_ping
):
"""Verify Redis < 7 fallback issues individual LPOPs and regroups correctly"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "6.2.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
# With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands
# Simulate: key1 has 2 values then None, key2 has 1 value then None
mock_pipeline.execute = AsyncMock(
return_value=[
b"val1",
b"val2",
None, # 3 LPOPs for key1
b"val3",
None, # 2 LPOPs for key2
]
)
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [
RedisPipelineLpopOperation(key="key1", count=3),
RedisPipelineLpopOperation(key="key2", count=2),
]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
assert len(results) == 2
assert results[0] == ["val1", "val2"] # 2 values, None filtered out
assert results[1] == ["val3"] # 1 value, None filtered out
# All 5 individual LPOPs should be queued, but only 1 execute() call
assert mock_pipeline.lpop.call_count == 5
mock_pipeline.execute.assert_called_once()
@pytest.mark.asyncio
async def test_async_rpush_pipeline_raises_on_per_command_error(
monkeypatch, redis_no_ping
):
"""Verify that per-command errors in pipeline results are raised, not silently dropped"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.rpush = MagicMock()
# Simulate: first RPUSH succeeds, second returns a per-command error
mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineRpushOperation
rpush_list = [
RedisPipelineRpushOperation(key="key1", values=["a"]),
RedisPipelineRpushOperation(key="key2", values=["b"]),
]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
with pytest.raises(Exception, match="WRONGTYPE"):
await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
@pytest.mark.asyncio
async def test_async_lpop_pipeline_raises_on_per_command_error(
monkeypatch, redis_no_ping
):
"""Verify that per-command errors in LPOP pipeline results are raised, not silently dropped"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "7.0.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
# Simulate: first LPOP succeeds, second returns a per-command error
mock_pipeline.execute = AsyncMock(return_value=[[b"val1"], Exception("WRONGTYPE")])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [
RedisPipelineLpopOperation(key="key1", count=10),
RedisPipelineLpopOperation(key="key2", count=10),
]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
with pytest.raises(Exception, match="WRONGTYPE"):
await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
@pytest.mark.asyncio
async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
"""Empty lpop_list should return empty list without touching Redis"""
@ -450,111 +130,6 @@ async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
mock_redis_instance.pipeline.assert_not_called()
@pytest.mark.asyncio
async def test_async_lpop_pipeline_propagates_redis_exception(
monkeypatch, redis_no_ping
):
"""Pipeline errors should propagate"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "7.0.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down"))
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)]
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
with pytest.raises(ConnectionError, match="Redis down"):
await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"redis_version",
[
# Standard cases
"7.0.0", # Standard Redis string version
7.0, # Valkey/ElastiCache float version (THE BUG this fix addresses)
7, # Integer version (e.g., from some Redis forks)
# Version < 7
"6", # String without dots, version < 7
# Malformed versions (fallback to 7)
"latest", # Non-numeric version
"", # Empty string
-7.0, # Negative float
# Format variations
" 7.0.0 ", # Whitespace (should be stripped)
"7.0.0-rc1", # Version with suffix
"10.0.0", # Double digit major version
],
)
async def test_async_lpop_with_float_redis_version(
monkeypatch, redis_no_ping, redis_version
):
"""
Test async_lpop with various Redis version formats (especially float).
This test specifically addresses the issue where AWS ElastiCache Valkey
returns redis_version as a float (e.g., 7.0) instead of a string (e.g., "7.0.0"),
which caused a 'float' object has no attribute 'split' error when trying to
use the Redis transaction buffer feature.
The fix converts the version to a string and handles edge cases like:
- Floats (7.0) and integers (7)
- Strings with/without dots ("7" vs "7.0.0")
- Malformed versions ("v7.0.0", "latest") - fallback to version 7
- Whitespace (" 7.0.0 ")
- Negative versions (fallback to version 7)
Related: Database deadlock issues when use_redis_transaction_buffer is enabled.
"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
# Create RedisCache instance
redis_cache = RedisCache()
redis_cache.redis_version = redis_version # Set the version to test
# Create an AsyncMock for the Redis client
mock_redis_instance = AsyncMock()
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
# Mock lpop to return a test value (Redis >= 7.0 behavior)
mock_redis_instance.lpop.return_value = [b"value1", b"value2"]
# Mock pipeline for Redis < 7.0 (used when major_version < 7)
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
# Make pipeline() a regular method (not async) that returns the mock
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
# Mock handle_lpop_count_for_older_redis_versions for Redis < 7
with patch.object(
redis_cache,
"handle_lpop_count_for_older_redis_versions",
return_value=[b"value1", b"value2"],
):
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
# Call async_lpop with count - this should not raise AttributeError
result = await redis_cache.async_lpop(key="test_key", count=2)
# Verify the method completed without error
assert result is not None
# LIT-3374: the namespace must be applied uniformly across every key-taking
# Redis operation, not just get/set/increment. Before the fix these paths wrote
# or read raw keys, so with a namespace configured the prefixed keys other

View file

@ -2889,3 +2889,76 @@ def test_streaming_chunks_share_one_chat_completion_id():
assert (
other_stream.chunk_parser(events[1]).id != ids[0]
), "a separate stream must get its own id, not a process-wide one"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"stream_options,expected_wire_stream_options",
[
({"include_usage": True, "include_obfuscation": False}, {"include_obfuscation": False}),
({"include_usage": True}, None),
],
)
async def test_acompletion_bridge_normalizes_stream_options_on_the_wire(
stream_options, expected_wire_stream_options
):
"""include_usage must be stripped from the /v1/responses body; include_obfuscation must survive as a dict."""
from unittest.mock import AsyncMock
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
responses_payload = {
"id": "resp_bridge_stream_options",
"object": "response",
"created_at": 1734366691,
"status": "completed",
"model": "gpt-5.5",
"output": [
{
"type": "message",
"id": "msg_1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"parallel_tool_calls": True,
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
"error": None,
"incomplete_details": None,
"instructions": None,
"metadata": None,
"temperature": None,
"tool_choice": "auto",
"tools": [],
"top_p": None,
"max_output_tokens": None,
"previous_response_id": None,
"reasoning": None,
"truncation": None,
"user": None,
}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.text = json.dumps(responses_payload)
mock_response.headers = httpx.Headers({})
mock_response.json.return_value = responses_payload
with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post:
mock_post.return_value = mock_response
await litellm.acompletion(
model="openai/responses/gpt-5.5",
messages=[{"role": "user", "content": "hi"}],
api_key="fake-api-key",
stream_options=stream_options,
)
mock_post.assert_called_once()
post_kwargs = mock_post.call_args.kwargs
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
if expected_wire_stream_options is None:
assert "stream_options" not in request_body
else:
assert request_body["stream_options"] == expected_wire_stream_options

View file

@ -29,7 +29,9 @@ from litellm.integrations.otel import ( # noqa: E402
from litellm.integrations.otel.plumbing import providers # noqa: E402
from litellm.integrations.otel.plumbing.context import ( # noqa: E402
reset_mcp_message_trace_carrier,
reset_mcp_message_transport_span_context,
set_mcp_message_trace_carrier,
set_mcp_message_transport_span_context,
set_request_root_span,
)
from litellm.integrations.otel.logger import OpenTelemetryV2 # noqa: E402
@ -56,9 +58,11 @@ def _reset_request_root_span():
_otel_context._request_root_span.set(None)
_otel_context._mcp_message_trace_carrier.set(None)
_otel_context._mcp_message_transport_span_context.set(None)
yield
_otel_context._request_root_span.set(None)
_otel_context._mcp_message_trace_carrier.set(None)
_otel_context._mcp_message_transport_span_context.set(None)
def _payload(**overrides):
@ -528,15 +532,15 @@ _MCP_SPAN_CASES = [
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_roots_and_links_transport_without_propagated_context(
def test_mcp_span_nests_under_transport_without_propagated_context(
make_payload, span_name
):
"""MCP and the HTTP transport are independent lifecycles (one streamable-HTTP
session multiplexes many messages), so per the MCP semconv the message span
must NOT nest under the session/transport span — that is what made it render
skewed at the session's start. With no propagated ``params._meta`` context it
starts its own root trace and records the transport span as a *link*, never
the parent."""
"""Almost no MCP client implements SEP-414, so ``params._meta`` normally carries
no trace context. Rooting the span there split one tool call into two traces
joined only by a link, which is how it surfaced in APM: the ``POST`` transaction
and the ``tools/call`` span shared no ``trace_id``. With no remote parent to
honor the span nests under the transport span instead, and records no link since
the transport is now the real parent."""
logger, exporter = _logger()
transport = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
@ -549,11 +553,75 @@ def test_mcp_span_roots_and_links_transport_without_propagated_context(
)
transport.end()
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
assert span.parent is not None
assert span.parent.span_id == transport.get_span_context().span_id
assert span.context.trace_id == transport.get_span_context().trace_id
assert span.links == ()
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_nests_under_this_messages_transport_not_the_session_opener(
make_payload, span_name
):
"""A *stateful* streamable-HTTP session runs every message on the single task
spawned by that session's ``initialize`` POST, so the ``_request_root_span``
ContextVar the ASGI request task writes is frozen at ``initialize`` inside the
handler and never sees the later ``tools/call`` POST. Nesting on that anchor
would hang every tool call of the session off the first request's (already
ended) span, rendering skewed at the session's start. The gateway resolves the
current message's transport on the request task and publishes it, so the span
parents to the POST that actually carried this message."""
logger, exporter = _logger()
session_opener = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
this_message = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
async def session_task():
token = set_mcp_message_transport_span_context(
this_message.get_span_context()
)
try:
await logger.async_log_success_event(
{"standard_logging_object": make_payload()}, None, None, None
)
finally:
reset_mcp_message_transport_span_context(token)
async def initialize_request():
# The anchor the session task inherits is the one ``initialize`` left behind;
# spawning here reproduces the SDK's session task, which outlives this request.
set_request_root_span(session_opener)
await asyncio.create_task(session_task())
asyncio.run(initialize_request())
session_opener.end()
this_message.end()
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
assert span.parent is not None
assert span.parent.span_id == this_message.get_span_context().span_id
assert span.context.trace_id == this_message.get_span_context().trace_id
assert span.parent.span_id != session_opener.get_span_context().span_id
assert span.context.trace_id != session_opener.get_span_context().trace_id
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_roots_without_transport_or_propagated_context(
make_payload, span_name
):
"""With neither a remote parent nor a transport span there is nothing to nest
under, so the span legitimately starts its own root trace with no links."""
logger, exporter = _logger()
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": make_payload()}, None, None, None
)
)
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
assert span.parent is None
assert span.context.trace_id != transport.get_span_context().trace_id
assert [link.context.span_id for link in span.links] == [
transport.get_span_context().span_id
]
assert span.links == ()
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
@ -641,10 +709,11 @@ def test_mcp_span_carries_authenticated_identity(make_payload, span_name):
assert span.attributes[LiteLLM.TEAM_ID] == "t1"
def test_mcp_span_malformed_traceparent_starts_root():
def test_mcp_span_malformed_traceparent_nests_under_transport():
"""A malformed traceparent in ``params._meta`` must not crash or parent to a
bogus span: the propagator ignores it, so the span starts its own root trace and
still links the transport span."""
bogus span: the propagator ignores it, leaving no remote parent, so the span
falls back to nesting under the transport span rather than starting a
disconnected root trace."""
logger, exporter = _logger()
transport = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
@ -661,9 +730,44 @@ def test_mcp_span_malformed_traceparent_starts_root():
reset_mcp_message_trace_carrier(token)
transport.end()
span = next(s for s in exporter.get_finished_spans() if s.name == "tools/list")
assert span.parent is None
assert span.parent is not None
assert span.parent.span_id == transport.get_span_context().span_id
assert span.links == ()
def test_mcp_span_links_this_messages_transport_when_context_is_propagated():
"""On the semconv path the transport is recorded as a link, and that link must
point at the POST carrying this message too. Reading the stale session anchor
would attribute the tool call to whichever request opened the session."""
logger, exporter = _logger()
session_opener = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
this_message = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
set_request_root_span(session_opener)
trace_token = set_mcp_message_trace_carrier(
{"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"}
)
transport_token = set_mcp_message_transport_span_context(
this_message.get_span_context()
)
try:
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": _mcp_list_payload()}, None, None, None
)
)
finally:
reset_mcp_message_transport_span_context(transport_token)
reset_mcp_message_trace_carrier(trace_token)
session_opener.end()
this_message.end()
span = next(s for s in exporter.get_finished_spans() if s.name == "tools/list")
assert span.parent is not None and span.parent.span_id == 0x2222222222222222
assert [link.context.span_id for link in span.links] == [
transport.get_span_context().span_id
this_message.get_span_context().span_id
]

View file

@ -1,10 +1,14 @@
import asyncio
from unittest.mock import AsyncMock
import pytest
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.proxy._types import CallTypes, UserAPIKeyAuth
from litellm.types.utils import GuardrailTracingDetail
from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailTracingDetail
class TestCustomGuardrailDeploymentHook:
@ -653,6 +657,53 @@ class TestGuardrailLoggingAggregation:
assert len(info) == 2
assert info[1]["guardrail_name"] == "test_guardrail"
def test_caller_metadata_does_not_divert_the_entry_from_the_reader(self):
"""A caller-supplied `metadata` field must not send the entry to a bucket the
spend log never reads. Routes in LITELLM_METADATA_ROUTES (/v1/messages,
/v1/responses, batches, files) seed `litellm_metadata`, and Claude Code sends
`metadata.user_id`, so both keys are present on the same request."""
request_data = {
"metadata": {"user_id": "device-account-session"},
"litellm_metadata": {"user_api_key_hash": "abc"},
}
self._invoke_add_log(request_data)
assert (
"standard_logging_guardrail_information" not in request_data["metadata"]
), "entry landed in the caller's metadata, where the spend log does not read it"
info = request_data["litellm_metadata"][
"standard_logging_guardrail_information"
]
assert len(info) == 1
assert info[0]["guardrail_name"] == "test_guardrail"
def test_entry_and_applied_guardrails_header_share_one_bucket(self):
"""The x-litellm-applied-guardrails writer and the guardrail-info writer must
resolve the same bucket, otherwise the response header and the spend log
disagree about whether the guardrail ran."""
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
request_data = {
"metadata": {"user_id": "device-account-session"},
"litellm_metadata": {},
}
self._invoke_add_log(request_data)
add_guardrail_to_applied_guardrails_header(
request_data=request_data, guardrail_name="test_guardrail"
)
buckets = {
key
for key in ("metadata", "litellm_metadata")
for field in ("standard_logging_guardrail_information", "applied_guardrails")
if field in request_data[key]
}
assert buckets == {"litellm_metadata"}
class TestGuardrailOtelSpanEmission:
"""Recording a guardrail emits its otel span inline, so every guardrail
@ -1394,6 +1445,38 @@ class TestEventTypeLogging:
assert len(logged_info) == 1
assert logged_info[0]["guardrail_status"] == "guardrail_intervened"
@pytest.mark.asyncio
async def test_log_guardrail_information_records_every_concurrent_guardrail(self):
"""Guardrails run concurrently (parallel pre_call/post_call, during_call) share one
request_data dict. Each must still record its own entry. The previous guard counted
entries in that shared dict, so a sibling's append made a guardrail think it had already
recorded and skip its own auto-record — silently dropping lifecycle logs the UI shows."""
from litellm.integrations.custom_guardrail import log_guardrail_information
from litellm.types.guardrails import GuardrailEventHooks
class SleeperGuardrail(CustomGuardrail):
def __init__(self, name, sleep):
super().__init__(guardrail_name=name, event_hook=GuardrailEventHooks.pre_call)
self._sleep = sleep
@log_guardrail_information
async def async_pre_call_hook(self, data: dict, **kwargs):
await asyncio.sleep(self._sleep)
return data
request_data = {"metadata": {}}
# Different sleeps guarantee overlapping execution windows: the faster guardrail
# records while the slower one is still awaiting, which is exactly what tripped the
# old shared-count guard.
await asyncio.gather(
SleeperGuardrail("guardrail-a", 0.05).async_pre_call_hook(data=request_data),
SleeperGuardrail("guardrail-b", 0.15).async_pre_call_hook(data=request_data),
)
logged = request_data["metadata"]["standard_logging_guardrail_information"]
assert {entry["guardrail_name"] for entry in logged} == {"guardrail-a", "guardrail-b"}
assert len(logged) == 2
def test_add_standard_logging_falls_back_to_event_hook_when_event_type_is_none(
self,
):
@ -1914,3 +1997,55 @@ class TestOnlyScanNewMessages:
cache.async_set_cache = AsyncMock(side_effect=RuntimeError("redis down"))
await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache)
def _guardrail_entries(request_data: dict) -> list:
container = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
entries = container.get("standard_logging_guardrail_information")
return entries if isinstance(entries, list) else []
class _NoopGuardrail(CustomGuardrail):
"""apply_guardrail that returns the inputs untouched and records nothing."""
@log_guardrail_information
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
return inputs
class _NoopSelfLoggingGuardrail(_NoopGuardrail):
records_own_guardrail_information = True
class TestRecordsOwnGuardrailInformation:
"""The @log_guardrail_information decorator must not synthesize an "allow"/"success"
entry for a no-op apply_guardrail when the guardrail sets
records_own_guardrail_information (LIT-4650)."""
@pytest.mark.asyncio
async def test_default_noop_apply_guardrail_is_auto_logged(self):
guardrail = _NoopGuardrail(guardrail_name="g1")
request_data: dict = {"model": "gpt-4o"}
await guardrail.apply_guardrail(
inputs=GenericGuardrailAPIInputs(texts=["x"]),
request_data=request_data,
input_type="request",
)
entries = _guardrail_entries(request_data)
assert len(entries) == 1
assert entries[0]["guardrail_status"] == "success"
@pytest.mark.asyncio
async def test_self_logging_noop_apply_guardrail_is_not_logged(self):
guardrail = _NoopSelfLoggingGuardrail(guardrail_name="g2")
request_data: dict = {"model": "gpt-4o"}
await guardrail.apply_guardrail(
inputs=GenericGuardrailAPIInputs(texts=["x"]),
request_data=request_data,
input_type="request",
)
assert _guardrail_entries(request_data) == []

View file

@ -59,8 +59,11 @@ def test_syncs_from_metadata_key():
assert result == [entry]
def test_metadata_wins_over_litellm_metadata():
"""metadata key takes precedence over litellm_metadata when both are present."""
def test_litellm_metadata_wins_over_caller_metadata():
"""When both keys are present the helper must read the bucket the writer used,
which get_or_create_metadata_bucket resolves to litellm_metadata. Reading the
caller's metadata instead is how a guardrail entry went missing from spend logs
on the routes that seed litellm_metadata."""
entry_meta = _make_slg_entry("from-metadata")
entry_lm = _make_slg_entry("from-litellm_metadata")
request_data = {
@ -74,7 +77,25 @@ def test_metadata_wins_over_litellm_metadata():
result = logging_obj.litellm_params["metadata"].get(
"standard_logging_guardrail_information"
)
assert result == [entry_meta]
assert result == [entry_lm]
def test_syncs_when_caller_sends_its_own_metadata():
"""The Claude Code shape: caller metadata present, guardrail entry in the seeded
litellm_metadata bucket. The entry must still reach the spend-log payload."""
entry = _make_slg_entry()
request_data = {
"metadata": {"user_id": "device-account-session"},
"litellm_metadata": {"standard_logging_guardrail_information": [entry]},
}
logging_obj = _FakeLogging()
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
result = logging_obj.litellm_params["metadata"].get(
"standard_logging_guardrail_information"
)
assert result == [entry]
def test_noop_when_no_guardrail_info():

View file

@ -279,6 +279,48 @@ class TestGuardrailSpanOnViolation(unittest.TestCase):
parent_span.context.span_id,
)
def test_post_call_failure_hook_emits_span_when_caller_sends_metadata(self):
"""On routes that seed ``litellm_metadata`` the guardrail entry lives there,
not in the caller's own ``metadata`` field. Reading a hard-coded ``metadata``
key drops the span for exactly the requests that carry both."""
otel, provider, exporter = _make_otel()
parent_span = provider.get_tracer(__name__).start_span(PROXY_SPAN_NAME)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test",
parent_otel_span=parent_span,
request_route="/v1/messages",
)
request_data = {
"model": "claude-haiku",
"messages": [{"role": "user", "content": "Hello"}],
"metadata": {"user_id": "device-account-session"},
"litellm_metadata": {
"standard_logging_guardrail_information": [
_slg_entry("guardrail_intervened", _bedrock_block_response())
],
},
}
_run(
otel.async_post_call_failure_hook(
request_data=request_data,
original_exception=Exception("guardrail blocked"),
user_api_key_dict=user_api_key_dict,
)
)
guardrail_spans = [
s for s in exporter.get_finished_spans() if s.name == GUARDRAIL_SPAN_NAME
]
self.assertEqual(
len(guardrail_spans),
1,
"the guardrail span must be emitted from the resolved metadata bucket, "
"not from a hard-coded 'metadata' key",
)
def test_handle_failure_and_post_call_failure_hook_dedupe(self):
"""When _handle_failure and async_post_call_failure_hook BOTH fire
for the same request (the production flow on a guardrail block),

View file

@ -4,12 +4,53 @@ import pytest
from litellm.litellm_core_utils.core_helpers import (
_FINISH_REASON_MAP,
get_or_create_metadata_bucket,
map_finish_reason,
reconstruct_model_name,
redact_nested_match_and_regex_keys,
)
class TestGetOrCreateMetadataBucket:
"""The single owner every guardrail writer and reader shares, so the response
header and the spend log can never disagree about which dict a record lives in."""
def test_prefers_litellm_metadata_when_both_present(self):
request_data = {"metadata": {"user_id": "caller"}, "litellm_metadata": {}}
key, bucket = get_or_create_metadata_bucket(request_data)
assert key == "litellm_metadata"
assert bucket is request_data["litellm_metadata"]
def test_uses_metadata_when_litellm_metadata_absent(self):
request_data = {"metadata": {"user_id": "caller"}}
key, bucket = get_or_create_metadata_bucket(request_data)
assert key == "metadata"
assert bucket is request_data["metadata"]
def test_creates_the_bucket_in_place_when_missing(self):
request_data: dict = {}
key, bucket = get_or_create_metadata_bucket(request_data)
assert key == "metadata"
assert request_data["metadata"] is bucket
bucket["k"] = "v"
assert request_data["metadata"]["k"] == "v"
def test_replaces_a_non_dict_bucket(self):
request_data = {"litellm_metadata": None}
key, bucket = get_or_create_metadata_bucket(request_data)
assert key == "litellm_metadata"
assert isinstance(request_data["litellm_metadata"], dict)
assert bucket is request_data["litellm_metadata"]
def test_reconstruct_model_name_prefers_deployment_value():
"""Ensure deployment metadata wins when reconstructing the model name."""

View file

@ -57,6 +57,98 @@ class MockDynamicGuardrail(CustomGuardrail):
return inputs
class MockRecordingGuardrail(CustomGuardrail):
"""Mock guardrail that records the request_data it was handed."""
def __init__(self, guardrail_name: str):
super().__init__(guardrail_name=guardrail_name)
self.request_data: Optional[dict] = None
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
self.request_data = request_data
return inputs
class TestAnthropicMessagesHandlerStreamingRequestData:
"""Post-call guardrails on streaming /v1/messages receive the response and identity metadata"""
@pytest.mark.asyncio
async def test_terminal_chunk_passes_assembled_response_and_metadata(self):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import Choices, Message, ModelResponse
handler = AnthropicMessagesHandler()
guardrail = MockRecordingGuardrail(guardrail_name="test")
mock_response = ModelResponse(
id="msg_123",
created=1234567890,
model="claude-sonnet-4-5",
object="chat.completion",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="Hello world", role="assistant"),
)
],
)
with (
patch.object(handler, "_check_streaming_has_ended", return_value=True),
patch(
"litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response",
return_value=mock_response,
),
):
await handler.process_output_streaming_response(
responses_so_far=[b"data: some chunk"],
guardrail_to_apply=guardrail,
litellm_logging_obj=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"),
request_data={"model": "claude-sonnet-4-5"},
)
assert guardrail.request_data is not None
assert guardrail.request_data["response"] is mock_response
assert (
guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1"
)
@pytest.mark.asyncio
async def test_mid_stream_chunk_passes_responses_so_far_and_metadata(self):
from litellm.proxy._types import UserAPIKeyAuth
handler = AnthropicMessagesHandler()
guardrail = MockRecordingGuardrail(guardrail_name="test")
responses_so_far = [b"data: some chunk"]
with (
patch.object(handler, "_check_streaming_has_ended", return_value=False),
patch.object(
handler, "get_streaming_string_so_far", return_value="partial text"
),
):
await handler.process_output_streaming_response(
responses_so_far=responses_so_far,
guardrail_to_apply=guardrail,
litellm_logging_obj=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(user_id="u-1", team_id="t-1"),
request_data={"model": "claude-sonnet-4-5"},
)
assert guardrail.request_data is not None
assert guardrail.request_data["responses"] is responses_so_far
assert (
guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1"
)
class TestAnthropicMessagesHandlerStreamingOutputProcessing:
"""Test streaming output processing functionality"""

View file

@ -2,6 +2,7 @@ import json
import os
import sys
import httpx
import pytest
from fastapi.testclient import TestClient
@ -9,6 +10,7 @@ sys.path.insert(0, os.path.abspath("../../../../.."))
from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm.anthropic_interface import messages
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.types.utils import Delta, ModelResponse, StreamingChoices
@ -37,6 +39,68 @@ def test_anthropic_experimental_pass_through_messages_handler():
assert mock_responses.call_args.kwargs["api_key"] == "test-api-key"
@pytest.mark.asyncio
async def test_openai_model_does_not_forward_stream_options_to_responses_api():
"""
Regression test for LIT-4779. `always_include_stream_usage` injects
stream_options={'include_usage': True} into every streaming request, but OpenAI
models on /v1/messages go to the Responses API, which 400s on that param.
"""
responses_payload = {
"id": "resp_stream_options",
"object": "response",
"created_at": 1734366691,
"status": "completed",
"model": "gpt-5.5",
"output": [
{
"type": "message",
"id": "msg_1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"parallel_tool_calls": True,
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
"error": None,
"incomplete_details": None,
"instructions": None,
"metadata": None,
"temperature": None,
"tool_choice": "auto",
"tools": [],
"top_p": None,
"max_output_tokens": None,
"previous_response_id": None,
"reasoning": None,
"truncation": None,
"user": None,
}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.text = json.dumps(responses_payload)
mock_response.headers = httpx.Headers({})
mock_response.json.return_value = responses_payload
with patch.object(AsyncHTTPHandler, "post", new_callable=AsyncMock) as mock_post:
mock_post.return_value = mock_response
await litellm.anthropic.messages.acreate(
max_tokens=100,
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="openai/gpt-5.5",
api_key="test-api-key",
stream_options={"include_usage": True},
)
mock_post.assert_called_once()
post_kwargs = mock_post.call_args.kwargs
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
assert "stream_options" not in request_body
def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_and_api_base_and_custom_values():
"""
Test that api key, api base, and extra kwargs are forwarded to litellm.completion for Azure models.

View file

@ -181,28 +181,6 @@ async def test_force_ipv4_transport():
litellm.disable_aiohttp_transport = original_disable
@pytest.mark.asyncio
async def test_ssl_context_transport():
"""Test transport creation with SSL context"""
# Create a test SSL context
ssl_context = ssl.create_default_context()
transport = AsyncHTTPHandler._create_async_transport(ssl_context=ssl_context)
assert transport is not None
try:
if isinstance(transport, LiteLLMAiohttpTransport):
# Get the client session and verify SSL context is passed through
client_session = transport._get_valid_client_session()
assert isinstance(client_session, ClientSession)
assert isinstance(client_session.connector, TCPConnector)
# Verify the connector has SSL context set by checking if it's using SSL
assert client_session.connector._ssl is not None
finally:
if isinstance(transport, LiteLLMAiohttpTransport):
await transport.aclose()
@pytest.mark.asyncio
async def test_aiohttp_disabled_transport():
"""Test transport creation with aiohttp disabled"""
@ -339,44 +317,6 @@ async def test_ssl_context_with_shared_session():
litellm.disable_aiohttp_transport = original_disable
@pytest.mark.asyncio
async def test_aiohttp_transport_trust_env_setting(monkeypatch):
"""Test that trust_env setting is properly configured in aiohttp transport"""
transports = []
try:
# Test 1: Default trust_env behavior
transport = AsyncHTTPHandler._create_aiohttp_transport()
transports.append(transport)
client_session = transport._get_valid_client_session()
# Default should be False (litellm.aiohttp_trust_env default)
default_trust_env = getattr(litellm, "aiohttp_trust_env", False)
assert client_session._trust_env == default_trust_env
# Test 2: Environment variable override
monkeypatch.setenv("AIOHTTP_TRUST_ENV", "True")
transport_with_env = AsyncHTTPHandler._create_aiohttp_transport()
transports.append(transport_with_env)
client_session_with_env = transport_with_env._get_valid_client_session()
# Should be True when environment variable is set
assert client_session_with_env._trust_env is True
# Test 3: Verify environment variable with False value
monkeypatch.setenv("AIOHTTP_TRUST_ENV", "False")
transport_with_false_env = AsyncHTTPHandler._create_aiohttp_transport()
transports.append(transport_with_false_env)
client_session_with_false_env = (
transport_with_false_env._get_valid_client_session()
)
# Should respect the litellm.aiohttp_trust_env setting when env var is False
assert client_session_with_false_env._trust_env == default_trust_env
finally:
for t in transports:
await t.aclose()
def test_get_ssl_configuration():
"""Test that get_ssl_configuration() returns a proper SSL context with certifi CA bundle
when no environment variables are set."""
@ -443,36 +383,6 @@ async def test_create_aiohttp_transport_with_shared_session():
assert not callable(transport.client) # Should not be callable
@pytest.mark.asyncio
async def test_create_aiohttp_transport_without_shared_session():
"""Test that _create_aiohttp_transport creates new session when none provided"""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
# Test without shared session
transport = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None)
# Verify the transport uses a lambda function (for backward compatibility)
assert callable(transport.client) # Should be a lambda function
@pytest.mark.asyncio
async def test_create_aiohttp_transport_with_closed_session():
"""Test that _create_aiohttp_transport creates new session when shared session is closed"""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
# Create a mock closed session
mock_session = MockClientSession()
mock_session.closed = True
# Test with closed session
transport = AsyncHTTPHandler._create_aiohttp_transport(
shared_session=mock_session # type: ignore
)
# Verify the transport creates a new session (lambda function)
assert callable(transport.client) # Should be a lambda function
@pytest.mark.asyncio
async def test_async_handler_with_shared_session():
"""Test AsyncHTTPHandler initialization with shared session"""
@ -622,27 +532,6 @@ async def test_session_reuse_integration():
await client2.close()
@pytest.mark.asyncio
async def test_session_validation():
"""Test that session validation works correctly"""
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
# Test with None session
transport1 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=None)
assert callable(transport1.client) # Should create lambda
# Test with closed session
mock_closed_session = MockClientSession()
mock_closed_session.closed = True
transport2 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_closed_session) # type: ignore
assert callable(transport2.client) # Should create lambda
# Test with valid session
mock_valid_session = MockClientSession()
transport3 = AsyncHTTPHandler._create_aiohttp_transport(shared_session=mock_valid_session) # type: ignore
assert transport3.client is mock_valid_session # Should reuse session
@pytest.mark.parametrize(
"env_curve,litellm_curve,expected_curve,should_call",
[

View file

@ -1,15 +1,10 @@
"""
Tests for LlmPassthroughRouteHandler and the guardrail_translation_mappings registry.
Tests for the guardrail_translation_mappings registry.
Validates:
- allm_passthrough_route is registered in the mappings (regression: this was the bug)
- Bedrock provider is dispatched to BedrockPassthroughGuardrailHandler
- Unknown provider skips apply_guardrail
"""
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.llms.pass_through.guardrail_translation import (
guardrail_translation_mappings,
)
@ -40,185 +35,3 @@ class TestRegistry:
is PassThroughEndpointHandler
)
def _make_guardrail() -> MagicMock:
g = MagicMock()
g.guardrail_name = "test-guard"
g.apply_guardrail = AsyncMock(return_value={"texts": []})
g.skip_system_message_in_guardrail = False
g.skip_tool_message_in_guardrail = False
return g
class TestLlmPassthroughRouteHandlerInput:
@pytest.mark.asyncio
async def test_bedrock_provider_delegates_to_bedrock_handler(self):
handler = LlmPassthroughRouteHandler()
data = {
"custom_llm_provider": "bedrock",
"endpoint": "model/anthropic.claude-3-sonnet/converse",
"data": {"messages": [{"role": "user", "content": [{"text": "hi"}]}]},
}
guardrail = _make_guardrail()
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
guardrail.apply_guardrail.assert_called_once()
@pytest.mark.asyncio
async def test_unknown_provider_skips_apply_guardrail(self):
handler = LlmPassthroughRouteHandler()
data = {
"custom_llm_provider": "some_unknown_provider",
"endpoint": "v1/chat/completions",
"data": {"messages": [{"role": "user", "content": "hi"}]},
}
guardrail = _make_guardrail()
result = await handler.process_input_messages(
data=data, guardrail_to_apply=guardrail
)
guardrail.apply_guardrail.assert_not_called()
assert result is data
@pytest.mark.asyncio
async def test_missing_provider_skips(self):
handler = LlmPassthroughRouteHandler()
data = {"endpoint": "foo/bar", "data": {}}
guardrail = _make_guardrail()
result = await handler.process_input_messages(
data=data, guardrail_to_apply=guardrail
)
guardrail.apply_guardrail.assert_not_called()
assert result is data
class TestLlmPassthroughRouteHandlerOutput:
@pytest.mark.asyncio
async def test_bedrock_provider_delegates_output_to_bedrock_handler(self):
handler = LlmPassthroughRouteHandler()
response = {
"output": {
"message": {
"role": "assistant",
"content": [{"text": "hello"}],
}
}
}
request_data = {
"custom_llm_provider": "bedrock",
"endpoint": "model/anthropic.claude-3-sonnet/converse",
}
guardrail = _make_guardrail()
await handler.process_output_response(
response=response,
guardrail_to_apply=guardrail,
request_data=request_data,
)
guardrail.apply_guardrail.assert_called_once()
@pytest.mark.asyncio
async def test_unknown_provider_skips_output(self):
handler = LlmPassthroughRouteHandler()
response = {"some": "response"}
request_data = {"custom_llm_provider": "unknown"}
guardrail = _make_guardrail()
result = await handler.process_output_response(
response=response,
guardrail_to_apply=guardrail,
request_data=request_data,
)
guardrail.apply_guardrail.assert_not_called()
assert result is response
class TestDeAnonymizeEventStream:
@pytest.mark.asyncio
async def test_bedrock_provider_dispatches_to_handler(self):
body = b"original-stream-bytes"
expected = b"de-anonymized-bytes"
proxy_logging_obj = MagicMock()
user_api_key_dict = MagicMock()
with patch(
"litellm.llms.bedrock.passthrough.guardrail_translation.handler."
"BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
new=AsyncMock(return_value=expected),
) as mock_handler:
result = await LlmPassthroughRouteHandler.de_anonymize_event_stream(
body_bytes=body,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
data={"custom_llm_provider": "bedrock"},
)
mock_handler.assert_awaited_once()
assert result == expected
@pytest.mark.asyncio
async def test_unknown_provider_returns_original_bytes(self):
body = b"original-stream-bytes"
result = await LlmPassthroughRouteHandler.de_anonymize_event_stream(
body_bytes=body,
proxy_logging_obj=MagicMock(),
user_api_key_dict=MagicMock(),
data={"custom_llm_provider": "anthropic"},
)
assert result is body
@pytest.mark.asyncio
async def test_missing_provider_returns_original_bytes(self):
body = b"original-stream-bytes"
result = await LlmPassthroughRouteHandler.de_anonymize_event_stream(
body_bytes=body,
proxy_logging_obj=MagicMock(),
user_api_key_dict=MagicMock(),
data={},
)
assert result is body
class TestSupportsEventStreamDeAnonymization:
def test_bedrock_converse_stream_is_supported(self):
assert (
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
"bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream"
)
is True
)
def test_bedrock_invoke_stream_is_not_supported(self):
assert (
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
"bedrock",
"model/us.amazon.nova-lite-v1:0/invoke-with-response-stream",
)
is False
)
def test_unknown_provider_is_not_supported(self):
assert (
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
"anthropic", "model/foo/converse-stream"
)
is False
)
def test_missing_provider_is_not_supported(self):
assert (
LlmPassthroughRouteHandler.supports_event_stream_de_anonymization(
None, "model/foo/converse-stream"
)
is False
)

View file

@ -226,6 +226,18 @@ def test_get_output_file_id_empty_output_info_falls_through_to_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_explicit_none_falls_through_to_output_config():
resp = {
"outputInfo": None,
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}},
}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_explicit_none_and_no_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response({"outputInfo": None}) == ""
def test_get_output_file_id_no_output_info_and_no_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response({}) == ""

View file

@ -63,6 +63,27 @@ def test_sso_descriptor_mapping_is_single_sourced():
)
def test_sso_descriptor_mapping_covers_saml_fields():
# SAML config is stored and read through the same descriptor table as the
# OAuth providers; the login path reads these env vars, so the save path must
# map every SAML field to its uppercase env var.
assert SSO_FIELD_ENV_VARS["saml_idp_metadata_url"] == "SAML_IDP_METADATA_URL"
assert SSO_FIELD_ENV_VARS["saml_idp_metadata_xml"] == "SAML_IDP_METADATA_XML"
assert SSO_FIELD_ENV_VARS["saml_sp_entity_id"] == "SAML_SP_ENTITY_ID"
assert SSO_FIELD_ENV_VARS["saml_allow_unsolicited"] == "SAML_ALLOW_UNSOLICITED"
def test_resolve_sso_config_resolves_saml_fields():
resolved = resolve_sso_config(
{"saml_idp_metadata_url": "https://idp.example.com/metadata"},
{"SAML_ALLOW_UNSOLICITED": "true"},
)
assert resolved.config.saml_idp_metadata_url == "https://idp.example.com/metadata"
assert resolved.provenance["saml_idp_metadata_url"] == "db"
assert resolved.config.saml_allow_unsolicited == "true"
assert resolved.provenance["saml_allow_unsolicited"] == "env"
def test_resolve_sso_config_returns_unmasked_secret_and_provenance():
# The resolver hands back plaintext; masking is the endpoint's job. If the
# resolver masked, the login path would consume a masked secret and fail.

View file

@ -729,10 +729,15 @@ async def test_openai_moderation_post_call_request_data_passthrough():
mock_make_request.assert_called_once()
# Guardrail info in the REAL request_data (not a throwaway)
guardrail_info_list = request_data["metadata"].get(
"standard_logging_guardrail_information"
# Guardrail info in the REAL request_data (not a throwaway). The unified hook
# seeds litellm_metadata, so read the bucket the resolver names rather than
# assuming "metadata"; the spend log reads it the same way.
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)]
guardrail_info_list = bucket.get("standard_logging_guardrail_information")
assert guardrail_info_list is not None
assert isinstance(guardrail_info_list[0]["guardrail_response"], dict)
assert "results" in guardrail_info_list[0]["guardrail_response"]

View file

@ -259,10 +259,15 @@ async def test_openai_moderation_streaming_end_of_stream_request_data_passthroug
):
pass
# Verify guardrail info reached the REAL request_data (not a throwaway)
guardrail_info_list = request_data["metadata"].get(
"standard_logging_guardrail_information"
# Verify guardrail info reached the REAL request_data (not a throwaway). The
# unified hook seeds litellm_metadata, so read the bucket the resolver names
# rather than assuming "metadata"; the spend log reads it the same way.
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
bucket = request_data[get_metadata_variable_name_from_kwargs(request_data)]
guardrail_info_list = bucket.get("standard_logging_guardrail_information")
assert (
guardrail_info_list is not None
), "Guardrail info should be in request_data after streaming"

View file

@ -114,6 +114,24 @@ def guardrail() -> HeadroomGuardrail:
return _make_guardrail()
def _recorded_guardrail_entries(request_data: dict) -> list:
for container_key in ("metadata", "litellm_metadata"):
container = request_data.get(container_key)
if isinstance(container, dict):
entries = container.get("standard_logging_guardrail_information")
if isinstance(entries, list):
return entries
return []
def _applied_guardrails(request_data: dict) -> list:
for container_key in ("metadata", "litellm_metadata"):
container = request_data.get(container_key)
if isinstance(container, dict) and isinstance(container.get("applied_guardrails"), list):
return container["applied_guardrails"]
return []
@pytest.mark.asyncio
async def test_apply_guardrail_compresses_and_returns_structured_messages(
guardrail: HeadroomGuardrail,
@ -123,6 +141,7 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages(
structured_messages=ORIGINAL_MESSAGES,
)
mock_response = _make_compress_response(COMPRESSED_MESSAGES)
request_data = {"model": "gpt-4o"}
with patch.object(
guardrail.async_handler,
@ -132,12 +151,19 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages(
):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"model": "gpt-4o"},
request_data=request_data,
input_type="request",
)
assert result.get("structured_messages") == COMPRESSED_MESSAGES
entries = _recorded_guardrail_entries(request_data)
assert len(entries) == 1
assert entries[0]["guardrail_name"] == "headroom"
assert entries[0]["guardrail_status"] == "success"
assert entries[0]["guardrail_provider"] == "headroom"
assert "headroom" in _applied_guardrails(request_data)
@pytest.mark.asyncio
async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present(
@ -719,6 +745,7 @@ async def test_apply_guardrail_bypass_header_skips_compression(
mock_post.assert_not_called()
assert result.get("structured_messages") == ORIGINAL_MESSAGES
assert _recorded_guardrail_entries(request_data) == []
@pytest.mark.asyncio
@ -729,16 +756,18 @@ async def test_apply_guardrail_response_type_passthrough(
texts=["some response text"],
structured_messages=ORIGINAL_MESSAGES,
)
request_data: dict = {"model": "gpt-4o"}
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
request_data=request_data,
input_type="response",
)
mock_post.assert_not_called()
assert result is inputs
assert _recorded_guardrail_entries(request_data) == []
@pytest.mark.asyncio
@ -746,16 +775,48 @@ async def test_apply_guardrail_empty_structured_messages_passthrough(
guardrail: HeadroomGuardrail,
):
inputs = GenericGuardrailAPIInputs(texts=["hello"])
request_data: dict = {"model": "gpt-4o"}
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
request_data=request_data,
input_type="request",
)
mock_post.assert_not_called()
assert result is inputs
assert _recorded_guardrail_entries(request_data) == []
assert "headroom" not in _applied_guardrails(request_data)
@pytest.mark.asyncio
async def test_passthrough_handler_does_not_log_headroom_as_run(
guardrail: HeadroomGuardrail,
):
"""Regression for LIT-4650.
A passthrough request drives headroom through PassThroughEndpointHandler, which
only supplies `texts` (no `structured_messages`). Headroom cannot compress that
shape and no-ops, so it must not appear in the spend log's
standard_logging_guardrail_information as a successful run.
"""
from litellm.llms.pass_through.guardrail_translation.handler import (
PassThroughEndpointHandler,
)
data = {"model": "gpt-4o", "messages": [{"role": "user", "content": "hello"}]}
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
await PassThroughEndpointHandler().process_input_messages(
data=data,
guardrail_to_apply=guardrail,
litellm_logging_obj=None,
)
mock_post.assert_not_called()
assert _recorded_guardrail_entries(data) == []
assert "headroom" not in _applied_guardrails(data)
@pytest.mark.asyncio
@ -884,6 +945,7 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed()
texts=["hello"],
structured_messages=ORIGINAL_MESSAGES,
)
request_data = {"model": "gpt-4o"}
with patch.object(
guardrail.async_handler,
@ -893,12 +955,18 @@ async def test_apply_guardrail_transport_error_fail_open_forwards_uncompressed()
):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
request_data=request_data,
input_type="request",
)
assert result["structured_messages"] == ORIGINAL_MESSAGES
entries = _recorded_guardrail_entries(request_data)
assert len(entries) == 1
assert entries[0]["guardrail_name"] == "headroom"
assert entries[0]["guardrail_status"] == "guardrail_failed_to_respond"
assert "headroom" in _applied_guardrails(request_data)
@pytest.mark.asyncio
async def test_apply_guardrail_http_error_fail_open_forwards_uncompressed():

View file

@ -3502,6 +3502,44 @@ async def test_single_scan_response_stays_a_dict():
assert isinstance(request_data["metadata"]["_model_armor_response"], dict)
@pytest.mark.asyncio
async def test_scan_result_reaches_the_logger_on_a_seeded_route():
"""On routes that seed `litellm_metadata` the scan result must land in that bucket
and be found by `_process_response`. Writing the file-scan result through the shared
resolver while the text-scan writers and the reader used a hard-coded `metadata` key
split the record in two, so the logged guardrail payload came back empty."""
guardrail = _make_guardrail()
pdf_b64 = base64.b64encode(PDF_BYTES).decode("utf-8")
request_data = {
"model": "claude-haiku",
"messages": [_file_message(pdf_b64)],
"metadata": {"user_id": "device-account-session"},
"litellm_metadata": {"guardrails": ["model-armor-test"]},
}
with patch.object(
guardrail.async_handler,
"post",
AsyncMock(return_value=_armor_response(blocked=False)),
):
await guardrail.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(),
cache=MagicMock(spec=DualCache),
data=request_data,
call_type="completion",
)
assert "_model_armor_response" not in request_data["metadata"]
assert "_model_armor_response" in request_data["litellm_metadata"]
before = len(request_data["litellm_metadata"].get("standard_logging_guardrail_information", []))
guardrail._process_response(response=None, request_data=request_data)
logged = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
assert len(logged) == before + 1
assert logged[-1]["guardrail_response"], "the logger recorded an empty Model Armor payload"
@pytest.mark.asyncio
async def test_pre_call_blocks_supported_document_with_undecodable_base64():
"""A supported document whose inline base64 will not decode cannot be scanned, so it fails closed."""

View file

@ -1,4 +1,5 @@
import json
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import httpx
@ -8,6 +9,9 @@ from litellm.exceptions import GuardrailRaisedException, ModifyResponseException
from litellm.proxy.guardrails.guardrail_hooks.straiker import initialize_guardrail
from litellm.proxy.guardrails.guardrail_hooks.straiker.straiker import (
StraikerGuardrail,
_build_usage,
_request_structured_messages,
_response_finish_reason,
)
from litellm.proxy.guardrails.guardrail_registry import (
guardrail_class_registry,
@ -17,7 +21,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
StraikerGuardrailConfigModel,
StraikerGuardrailConfigModelOptionalParams,
)
from litellm.types.utils import Choices, Message, ModelResponse, Usage
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Choices,
Function,
Message,
ModelResponse,
Usage,
)
def _mock_response(action: str, turn_id: str = "turn-1", schema_version: str = "1", **extra) -> MagicMock:
@ -208,7 +219,9 @@ async def test_request_envelope_transport_and_shape():
"metadata": {"user_api_key_alias": "team-key", "agent_id": "chatbot-app", "app_name": "Chatbot"},
}
out = await g.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj())
out = await g.apply_guardrail(
inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj()
)
assert out is inputs
url = g.async_handler.post.call_args.args[0]
@ -232,6 +245,35 @@ async def test_request_envelope_transport_and_shape():
assert "metadata" not in payload
@pytest.mark.asyncio
async def test_request_envelope_ignores_unsupported_opaque_items():
g = _make_guardrail()
g.async_handler.post.return_value = _mock_response("NONE")
await g.apply_guardrail(
inputs={
"texts": ["hello"],
"tools": [
object(),
{
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
},
],
},
request_data={"model": "m", "messages": [{"role": "user", "content": "hello"}]},
input_type="request",
logging_obj=_logging_obj(),
)
assert _posted_payload(g)["request"]["tools"] == [
{
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
}
]
@pytest.mark.asyncio
async def test_webhook_metadata_session_id_and_opaque_passthrough():
g = _make_guardrail()
@ -331,6 +373,64 @@ async def test_context_session_id_from_request_metadata():
assert "metadata" not in payload
@pytest.mark.asyncio
async def test_context_mode_from_string_event_hook():
g = _make_guardrail(event_hook="pre_call")
g.async_handler.post.return_value = _mock_response("NONE")
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"model": "m"},
input_type="request",
logging_obj=_logging_obj(),
)
assert _posted_payload(g)["context"]["mode"] == ["pre_call"]
@pytest.mark.asyncio
async def test_context_mode_from_list_event_hook():
from litellm.types.guardrails import GuardrailEventHooks
g = _make_guardrail(event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call])
g.async_handler.post.return_value = _mock_response("NONE")
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"model": "m"},
input_type="request",
logging_obj=_logging_obj(),
)
assert _posted_payload(g)["context"]["mode"] == ["pre_call", "post_call"]
@pytest.mark.asyncio
async def test_context_mode_from_tagged_mode_is_flattened_and_deduped():
from litellm.types.guardrails import Mode
g = _make_guardrail(
event_hook=Mode(tags={"team-a": "pre_call", "team-b": ["post_call", "pre_call"]}, default="post_call")
)
g.async_handler.post.return_value = _mock_response("NONE")
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"model": "m"},
input_type="request",
logging_obj=_logging_obj(),
)
assert _posted_payload(g)["context"]["mode"] == ["post_call", "pre_call"]
@pytest.mark.asyncio
async def test_context_mode_omitted_when_event_hook_absent():
g = _make_guardrail(event_hook=None)
g.async_handler.post.return_value = _mock_response("NONE")
await g.apply_guardrail(
inputs={"texts": ["x"]},
request_data={"model": "m"},
input_type="request",
logging_obj=_logging_obj(),
)
assert "mode" not in _posted_payload(g)["context"]
@pytest.mark.asyncio
async def test_identity_key_and_team_coalesce_alias_over_id():
g = _make_guardrail()
@ -426,6 +526,7 @@ async def test_application_source_from_agent_id():
)
assert _posted_payload(g)["application"] == {"source": "analytics-app", "name": "Analytics"}
@pytest.mark.asyncio
async def test_request_block_raises_guardrail_exception_with_reason():
g = _make_guardrail()
@ -560,6 +661,34 @@ async def test_response_envelope_and_block_replaces_response():
assert payload["request"]["structured_messages"] == [{"role": "user", "content": "original prompt"}]
@pytest.mark.asyncio
async def test_post_call_resolves_request_from_responses_input_when_messages_absent():
g = _make_guardrail()
g.async_handler.post.return_value = _mock_response("NONE")
response = ModelResponse(
choices=[Choices(finish_reason="stop", index=0, message=Message(content="answer", role="assistant"))],
model="gpt-4o-mini",
)
request_data = {
"model": "gpt-4o-mini",
"input": "responses-surface prompt",
"response": response,
"litellm_metadata": {"user_api_key_request_route": "/v1/responses"},
}
await g.apply_guardrail(
inputs={"texts": ["answer"], "model": "gpt-4o-mini"},
request_data=request_data,
input_type="response",
logging_obj=_logging_obj(),
)
payload = _posted_payload(g)
assert payload["event"]["type"] == "post_call"
messages = payload["request"]["structured_messages"]
assert any(m.get("content") == "responses-surface prompt" for m in messages)
@pytest.mark.asyncio
async def test_post_call_fail_closed_raises_modify_response_exception():
g = _make_guardrail(unreachable_fallback="fail_closed")
@ -731,3 +860,210 @@ async def test_unreachable_http_status_fail_closed_blocks():
await g.apply_guardrail(
inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
)
@pytest.mark.asyncio
async def test_post_call_preserves_anthropic_tool_blocks_in_request_messages():
g = _make_guardrail()
g.async_handler.post.return_value = _mock_response("NONE")
anthropic_messages = [
{"role": "user", "content": "What's the weather in Paris?"},
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_1",
"name": "get_weather",
"input": {"city": "Paris"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": "18C, cloudy",
}
],
},
]
response = {
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "Mild and cloudy."}],
"stop_reason": "end_turn",
"model": "claude-sonnet-5",
}
await g.apply_guardrail(
inputs={"texts": ["Mild and cloudy."], "model": "claude-sonnet-5"},
request_data={
"model": "claude-sonnet-5",
"messages": anthropic_messages,
"response": response,
},
input_type="response",
logging_obj=_logging_obj(),
)
payload = _posted_payload(g)
assert payload["request"]["structured_messages"] == anthropic_messages
assert payload["response"]["finish_reason"] == "end_turn"
@pytest.mark.asyncio
async def test_pre_call_preserves_anthropic_tool_blocks_in_structured_messages():
g = _make_guardrail()
g.async_handler.post.return_value = _mock_response("NONE")
anthropic_messages = [
{
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": "toolu_1",
"name": "get_weather",
"input": {"city": "Paris"},
}
],
},
{
"role": "user",
"content": [
{
"type": "tool_result",
"tool_use_id": "toolu_1",
"content": "18C",
}
],
},
]
await g.apply_guardrail(
inputs={"structured_messages": anthropic_messages, "model": "claude-sonnet-5"},
request_data={"model": "claude-sonnet-5", "messages": anthropic_messages},
input_type="request",
logging_obj=_logging_obj(),
)
assert _posted_payload(g)["request"]["structured_messages"] == anthropic_messages
@pytest.mark.asyncio
async def test_response_finish_reason_from_openai_choices_still_works():
g = _make_guardrail()
g.async_handler.post.return_value = _mock_response("NONE")
response = ModelResponse(
choices=[Choices(finish_reason="tool_calls", index=0, message=Message(content=None, role="assistant"))],
model="gpt-4o-mini",
)
await g.apply_guardrail(
inputs={
"texts": [],
"tool_calls": [
ChatCompletionMessageToolCall(
id="c1",
type="function",
function=Function(name="f", arguments="{}"),
)
],
},
request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}], "response": response},
input_type="response",
logging_obj=_logging_obj(),
)
payload = _posted_payload(g)
assert payload["response"]["finish_reason"] == "tool_calls"
assert payload["response"]["tool_calls"] == [
{"id": "c1", "type": "function", "function": {"name": "f", "arguments": "{}"}}
]
@pytest.mark.parametrize(
("response", "expected"),
[
(None, None),
({"choices": "invalid"}, None),
({"choices": [{"finish_reason": "length"}]}, "length"),
({"choices": [{"stop_reason": "end_turn"}]}, "end_turn"),
({"choices": [{}]}, None),
(SimpleNamespace(stop_reason="end_turn"), "end_turn"),
],
)
def test_response_finish_reason_handles_supported_shapes(response, expected):
assert _response_finish_reason(response) == expected
@pytest.mark.parametrize(
"request_data",
[
{"input": ["ssn 123-45-6789"], "litellm_metadata": {"user_api_key_request_route": "/vllm/v1/embeddings"}},
{"input": [[1, 2, 3]], "litellm_metadata": {}},
{"input": "confidential memo", "litellm_metadata": {}},
{"input": "confidential memo"},
],
)
def test_request_messages_not_resolved_for_unmapped_surfaces(request_data):
"""Bodies from surfaces without a translation handler yield no messages, and never raise."""
assert _request_structured_messages(request_data) is None
@pytest.mark.parametrize(
("request_data", "expected"),
[
(
{"messages": [{"role": "user", "content": "hi"}], "litellm_metadata": {}},
[{"role": "user", "content": "hi"}],
),
(
{
"input": [{"role": "user", "content": "weather in Paris?"}],
"litellm_metadata": {"user_api_key_request_route": "/v1/responses"},
},
[{"role": "user", "content": "weather in Paris?"}],
),
],
)
def test_request_messages_resolved_for_mapped_surfaces(request_data, expected):
assert _request_structured_messages(request_data) == expected
@pytest.mark.parametrize(
("response", "expected"),
[
({"usage": {"input_tokens": 10, "output_tokens": 5}}, (10, 5)),
({"usage": {"prompt_tokens": 7, "completion_tokens": 3}}, (7, 3)),
(SimpleNamespace(usage=Usage(prompt_tokens=7, completion_tokens=3)), (7, 3)),
({"usage": {"prompt_tokens": 0, "input_tokens": 99}}, (0, None)),
({"usage": {}}, None),
({}, None),
],
)
def test_build_usage_handles_openai_and_anthropic_shapes(response, expected):
usage = _build_usage(response)
if expected is None:
assert usage is None
else:
assert (usage.input_tokens, usage.output_tokens) == expected
@pytest.mark.asyncio
async def test_anthropic_non_streaming_response_reports_usage():
g = _make_guardrail()
g.async_handler.post.return_value = _mock_response("NONE")
await g.apply_guardrail(
inputs={"texts": ["hello"]},
request_data={
"model": "claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hi"}],
"response": {
"stop_reason": "end_turn",
"usage": {"input_tokens": 10, "output_tokens": 5},
},
},
input_type="response",
logging_obj=_logging_obj(),
)
payload = _posted_payload(g)
assert payload["usage"] == {"input_tokens": 10, "output_tokens": 5}
assert payload["response"]["finish_reason"] == "end_turn"

View file

@ -4,7 +4,10 @@ import pytest
import litellm
from litellm.caching import DualCache
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_skip_system_message_for_guardrail,
@ -1490,3 +1493,109 @@ class TestStreamingTransform:
# None holdback treated as 0: full text emitted, no crash.
assert "".join(_delta_text(i) for i in out) == "ABCDEF"
def _applied_guardrails(data: dict) -> list:
for key in ("metadata", "litellm_metadata"):
meta = data.get(key)
if isinstance(meta, dict) and isinstance(meta.get("applied_guardrails"), list):
return meta["applied_guardrails"]
return []
class _TextsOnlyTranslation(BaseTranslation):
"""Mimics a passthrough handler: hands the guardrail only `texts`, never
structured_messages, so a structured_messages-based guardrail no-ops."""
async def process_input_messages(self, data, guardrail_to_apply, litellm_logging_obj=None): # type: ignore[override]
await guardrail_to_apply.apply_guardrail(
inputs={"texts": ["payload"]},
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
return data
async def process_output_response( # type: ignore[override]
self,
response,
guardrail_to_apply,
litellm_logging_obj=None,
user_api_key_dict=None,
request_data=None,
):
return response
class _SelfLoggingGuardrail(CustomGuardrail):
records_own_guardrail_information = True
def __init__(self, *, self_add: bool):
super().__init__(guardrail_name="self-logging")
self._self_add = self_add
def should_run_guardrail(self, data, event_type): # type: ignore[override]
return True
@log_guardrail_information
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
if self._self_add:
from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
return inputs
class _AutoLoggingGuardrail(CustomGuardrail):
def __init__(self):
super().__init__(guardrail_name="auto-logging")
def should_run_guardrail(self, data, event_type): # type: ignore[override]
return True
@log_guardrail_information
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
return inputs
class TestAppliedGuardrailsReflectsExecution:
"""The unified hook must not auto-mark a self-logging guardrail
(records_own_guardrail_information) as applied; such a guardrail owns that
decision and marks itself only when it actually ran (LIT-4650). Ordinary
guardrails are still auto-marked by the hook after dispatch."""
@staticmethod
def _data(guardrail):
return {
"guardrail_to_apply": guardrail,
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello world"}],
}
async def _run(self, guardrail):
unified_module.endpoint_guardrail_translation_mappings = {CallTypes.pass_through: _TextsOnlyTranslation}
data = self._data(guardrail)
await UnifiedLLMGuardrails().async_pre_call_hook(
user_api_key_dict=None,
cache=DualCache(),
data=data,
call_type=CallTypes.pass_through.value,
)
return data
@pytest.mark.asyncio
async def test_self_logging_guardrail_is_not_auto_marked_applied(self):
data = await self._run(_SelfLoggingGuardrail(self_add=False))
assert "self-logging" not in _applied_guardrails(data)
@pytest.mark.asyncio
async def test_self_logging_guardrail_that_self_marks_is_applied(self):
data = await self._run(_SelfLoggingGuardrail(self_add=True))
assert "self-logging" in _applied_guardrails(data)
@pytest.mark.asyncio
async def test_ordinary_guardrail_is_auto_marked_applied(self):
data = await self._run(_AutoLoggingGuardrail())
assert "auto-logging" in _applied_guardrails(data)

View file

@ -1,3 +1,5 @@
import pytest
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.guardrails.guardrail_registry import (
get_guardrail_initializer_from_hooks,
@ -32,6 +34,44 @@ def test_noma_registry_resolution():
assert "noma_v2" in guardrail_initializer_registry
@pytest.mark.parametrize(
"configured, expected",
[(None, True), (False, False), (True, True)],
)
def test_initialize_guardrail_run_in_parallel_preserves_constructor_default(configured, expected):
"""
A guardrail whose constructor sets run_in_parallel=True must keep that default when
the config omits the key; only an explicit config value may override it. The
previous code wrote bool(None)==False on every instance, silently disabling the
opt-in for such guardrails.
"""
from litellm.proxy.guardrails import guardrail_registry as registry_module
def _initializer(litellm_params, guardrail):
return CustomGuardrail(
guardrail_name=guardrail["guardrail_name"],
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
run_in_parallel=True,
)
registry_module.guardrail_initializer_registry["parallel_default_test"] = _initializer
try:
params = {"guardrail": "parallel_default_test", "mode": "pre_call"}
if configured is not None:
params["run_in_parallel"] = configured
handler = InMemoryGuardrailHandler()
result = handler.initialize_guardrail(
guardrail={"guardrail_name": "cf-parallel-default", "litellm_params": params},
)
stored = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
assert stored.run_in_parallel is expected
finally:
registry_module.guardrail_initializer_registry.pop("parallel_default_test", None)
def test_update_in_memory_guardrail():
handler = InMemoryGuardrailHandler()
handler.guardrail_id_to_custom_guardrail["123"] = CustomGuardrail(

View file

@ -62,3 +62,27 @@ def test_initialize_guardrail_preserves_guardrail_info():
assert result["guardrail_info"] == {"type": "PII", "description": "masks PII"}
stored = guardrail_handler.IN_MEMORY_GUARDRAILS[result["guardrail_id"]]
assert stored["guardrail_info"] == {"type": "PII", "description": "masks PII"}
@pytest.mark.parametrize(
"config_value, expected",
[(True, True), (False, False), (None, False)],
)
def test_initialize_guardrail_sets_run_in_parallel(config_value, expected):
"""run_in_parallel from litellm_params must reach the built guardrail instance."""
litellm_params = {
"guardrail": SupportedGuardrailIntegrations.PRESIDIO.value,
"mode": "pre_call",
"presidio_analyzer_api_base": "https://fakelink.com/v1/presidio/analyze",
"presidio_anonymizer_api_base": "https://fakelink.com/v1/presidio/anonymize",
}
if config_value is not None:
litellm_params["run_in_parallel"] = config_value
guardrail_handler = InMemoryGuardrailHandler()
result = guardrail_handler.initialize_guardrail(
guardrail={"guardrail_name": "test_parallel_flag", "litellm_params": litellm_params},
)
custom_guardrail = guardrail_handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
assert custom_guardrail.run_in_parallel is expected

View file

@ -0,0 +1,644 @@
"""
Regression tests for SAML 2.0 SSO (SP- and IdP-initiated) on the admin UI.
These exercise the real OneLogin python3-saml validation by generating signed
SAML responses with a freshly minted IdP keypair, so a mutation that weakens
signature, signing-requirement, expiry, replay or attribute-mapping handling
makes a test fail.
"""
import base64
import datetime
import time
import pytest
from fastapi import HTTPException, Request
pytest.importorskip(
"onelogin", reason="python3-saml (saml extra) is required for SAML SSO tests"
)
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID
from onelogin.saml2.utils import OneLogin_Saml2_Utils
from starlette.datastructures import URL
from typing import cast
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.management_endpoints.sso.saml_sso import (
_SAML_AUTHN_REQUEST_CACHE_PREFIX,
_SAML_AUTHN_STATE_COOKIE,
_SAML_MAX_POST_BYTES,
_SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS,
_SAML_REPLAY_GUARD_MAX_TTL_SECONDS,
SAMLAuthHandler,
)
def _shared_cache(store=None):
"""A DualCache whose replay guard is backed by a shared, atomic store.
An InMemoryCache instance stands in for Redis; passing the same instance to
two DualCaches simulates two workers sharing one atomic backend."""
return DualCache(redis_cache=cast(RedisCache, store or InMemoryCache()))
IDP_ENTITY = "https://idp.example.com/metadata"
SP_ENTITY = "https://proxy.example.com/sso/saml/metadata"
ACS = "https://proxy.example.com/sso/saml/callback"
SSO_URL = "https://idp.example.com/sso"
PROXY_BASE_URL = "https://proxy.example.com"
def _make_idp_keypair():
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "idp.example.com")])
cert = (
x509.CertificateBuilder()
.subject_name(name)
.issuer_name(name)
.public_key(key.public_key())
.serial_number(x509.random_serial_number())
.not_valid_before(datetime.datetime.utcnow() - datetime.timedelta(days=1))
.not_valid_after(datetime.datetime.utcnow() + datetime.timedelta(days=365))
.sign(key, hashes.SHA256())
)
key_pem = key.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.TraditionalOpenSSL,
serialization.NoEncryption(),
).decode()
cert_pem = cert.public_bytes(serialization.Encoding.PEM).decode()
return key_pem, cert_pem
def _idp_metadata_xml(cert_pem):
cert_body = "".join(
line for line in cert_pem.splitlines() if "CERTIFICATE" not in line
)
return (
'<?xml version="1.0"?>'
f'<EntityDescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" entityID="{IDP_ENTITY}">'
'<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">'
'<KeyDescriptor use="signing"><KeyInfo xmlns="http://www.w3.org/2000/09/xmldsig#">'
f"<X509Data><X509Certificate>{cert_body}</X509Certificate></X509Data>"
"</KeyInfo></KeyDescriptor>"
'<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" '
f'Location="{SSO_URL}"/>'
"</IDPSSODescriptor></EntityDescriptor>"
)
def _saml_time(delta_seconds):
t = datetime.datetime.utcnow() + datetime.timedelta(seconds=delta_seconds)
return t.strftime("%Y-%m-%dT%H:%M:%SZ")
def _build_signed_response(
key_pem,
cert_pem,
*,
in_response_to=None,
response_level_in_response_to=True,
email="alice@example.com",
attributes=None,
not_before_delta=-60,
not_on_or_after_delta=300,
sign=True,
):
if attributes is None:
attributes = {
"email": [email],
"givenName": ["Alice"],
"sn": ["Smith"],
"role": ["internal_user"],
}
assertion_id = "_assertion_" + OneLogin_Saml2_Utils.generate_unique_id()
response_id = "_response_" + OneLogin_Saml2_Utils.generate_unique_id()
not_before = _saml_time(not_before_delta)
not_on_or_after = _saml_time(not_on_or_after_delta)
issue_instant = _saml_time(-1)
irt = f'InResponseTo="{in_response_to}"' if in_response_to else ""
response_irt = irt if response_level_in_response_to else ""
attr_xml = "".join(
f'<saml:Attribute Name="{name}">'
+ "".join(f"<saml:AttributeValue>{v}</saml:AttributeValue>" for v in values)
+ "</saml:Attribute>"
for name, values in attributes.items()
)
assertion = (
'<saml:Assertion xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion" '
f'ID="{assertion_id}" Version="2.0" IssueInstant="{issue_instant}">'
f"<saml:Issuer>{IDP_ENTITY}</saml:Issuer>"
"<saml:Subject>"
'<saml:NameID Format="urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress">'
f"{email}</saml:NameID>"
'<saml:SubjectConfirmation Method="urn:oasis:names:tc:SAML:2.0:cm:bearer">'
f'<saml:SubjectConfirmationData {irt} NotOnOrAfter="{not_on_or_after}" Recipient="{ACS}"/>'
"</saml:SubjectConfirmation></saml:Subject>"
f'<saml:Conditions NotBefore="{not_before}" NotOnOrAfter="{not_on_or_after}">'
f"<saml:AudienceRestriction><saml:Audience>{SP_ENTITY}</saml:Audience>"
"</saml:AudienceRestriction></saml:Conditions>"
f'<saml:AuthnStatement AuthnInstant="{issue_instant}" SessionIndex="_session">'
"<saml:AuthnContext><saml:AuthnContextClassRef>"
"urn:oasis:names:tc:SAML:2.0:ac:classes:Password"
"</saml:AuthnContextClassRef></saml:AuthnContext></saml:AuthnStatement>"
f"<saml:AttributeStatement>{attr_xml}</saml:AttributeStatement>"
"</saml:Assertion>"
)
if sign:
signed = OneLogin_Saml2_Utils.add_sign(assertion, key_pem, cert_pem)
assertion = (signed.decode() if isinstance(signed, bytes) else signed).replace(
'<?xml version="1.0"?>', ""
)
return (
'<?xml version="1.0"?>'
'<samlp:Response xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol" '
'xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion" '
f'ID="{response_id}" Version="2.0" IssueInstant="{issue_instant}" '
f'Destination="{ACS}" {response_irt}>'
f"<saml:Issuer>{IDP_ENTITY}</saml:Issuer>"
'<samlp:Status><samlp:StatusCode Value="urn:oasis:names:tc:SAML:2.0:status:Success"/>'
"</samlp:Status>"
f"{assertion}</samlp:Response>"
)
def _b64(xml):
return base64.b64encode(xml.encode()).decode()
def _fake_request(cookies=None):
return type(
"Req",
(),
{
"base_url": URL(PROXY_BASE_URL + "/"),
"query_params": {},
"cookies": cookies or {},
},
)()
async def _acs(b64, cache, cookies=None):
return await SAMLAuthHandler.handle_acs(
_fake_request(cookies), cache, {"SAMLResponse": b64}
)
@pytest.fixture
def saml_env(monkeypatch):
key_pem, cert_pem = _make_idp_keypair()
monkeypatch.setenv("SAML_IDP_METADATA_XML", _idp_metadata_xml(cert_pem))
monkeypatch.setenv("SAML_SP_ENTITY_ID", SP_ENTITY)
monkeypatch.setenv("PROXY_BASE_URL", PROXY_BASE_URL)
for var in (
"SAML_IDP_METADATA_URL",
"SAML_ATTRIBUTE_EMAIL",
"SAML_ATTRIBUTE_TEAM_IDS",
"SAML_ALLOW_UNSOLICITED",
"ALLOWED_EMAIL_DOMAINS",
):
monkeypatch.delenv(var, raising=False)
return key_pem, cert_pem
@pytest.fixture
def saml_env_idp_initiated(saml_env, monkeypatch):
monkeypatch.setenv("SAML_ALLOW_UNSOLICITED", "true")
return saml_env
@pytest.mark.asyncio
async def test_valid_idp_initiated_login_maps_assertion_to_user(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(key_pem, cert_pem)
result = await _acs(_b64(resp), _shared_cache())
assert result.email == "alice@example.com"
assert result.id == "alice@example.com"
assert result.first_name == "Alice"
assert result.last_name == "Smith"
assert result.user_role == LitellmUserRoles.INTERNAL_USER
assert result.provider == "saml"
@pytest.mark.asyncio
async def test_tampered_assertion_is_rejected(saml_env):
key_pem, cert_pem = saml_env
resp = _build_signed_response(key_pem, cert_pem)
tampered = resp.replace("alice@example.com", "attacker@example.com")
with pytest.raises(HTTPException) as exc:
await _acs(_b64(tampered), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_unsigned_assertion_is_rejected(saml_env):
key_pem, cert_pem = saml_env
resp = _build_signed_response(key_pem, cert_pem, sign=False)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_signature_from_untrusted_key_is_rejected(saml_env):
_, cert_pem = saml_env
attacker_key, attacker_cert = _make_idp_keypair()
resp = _build_signed_response(attacker_key, attacker_cert)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_expired_assertion_is_rejected(saml_env):
key_pem, cert_pem = saml_env
resp = _build_signed_response(
key_pem, cert_pem, not_before_delta=-7200, not_on_or_after_delta=-3600
)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_sp_initiated_unknown_in_response_to_is_rejected(saml_env):
key_pem, cert_pem = saml_env
resp = _build_signed_response(key_pem, cert_pem, in_response_to="_never_issued")
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_sp_initiated_known_request_succeeds_once_then_replay_rejected(saml_env):
key_pem, cert_pem = saml_env
cache = DualCache()
request_id = "_authn_req_known"
cache.set_cache(
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
)
resp = _build_signed_response(key_pem, cert_pem, in_response_to=request_id)
cookies = {_SAML_AUTHN_STATE_COOKIE: request_id}
result = await _acs(_b64(resp), cache, cookies=cookies)
assert result.email == "alice@example.com"
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), cache, cookies=cookies)
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_sp_initiated_response_not_bound_to_browser_is_rejected(saml_env):
key_pem, cert_pem = saml_env
cache = DualCache()
request_id = "_authn_req_known"
cache.set_cache(
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
)
resp = _build_signed_response(key_pem, cert_pem, in_response_to=request_id)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), cache)
assert exc.value.status_code == 401
with pytest.raises(HTTPException) as exc:
await _acs(
_b64(resp), cache, cookies={_SAML_AUTHN_STATE_COOKIE: "_attacker_request"}
)
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_subjectconfirmation_only_in_response_to_without_cookie_is_rejected(
saml_env_idp_initiated,
):
"""An IdP that stamps InResponseTo only on the SubjectConfirmationData (not the
Response element) is still solicited and must be browser-bound: with unsolicited
explicitly allowed, a missing cookie must still 401 rather than slip through."""
key_pem, cert_pem = saml_env_idp_initiated
cache = DualCache()
request_id = "_authn_req_known"
cache.set_cache(
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
)
resp = _build_signed_response(
key_pem,
cert_pem,
in_response_to=request_id,
response_level_in_response_to=False,
)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), cache)
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_subjectconfirmation_only_in_response_to_with_cookie_succeeds(saml_env):
key_pem, cert_pem = saml_env
cache = DualCache()
request_id = "_authn_req_known"
cache.set_cache(
key=f"{_SAML_AUTHN_REQUEST_CACHE_PREFIX}:{request_id}", value="1", ttl=600
)
resp = _build_signed_response(
key_pem,
cert_pem,
in_response_to=request_id,
response_level_in_response_to=False,
)
result = await _acs(
_b64(resp), cache, cookies={_SAML_AUTHN_STATE_COOKIE: request_id}
)
assert result.email == "alice@example.com"
@pytest.mark.asyncio
async def test_unsolicited_response_rejected_by_default(saml_env):
key_pem, cert_pem = saml_env
resp = _build_signed_response(key_pem, cert_pem)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_idp_initiated_assertion_replay_is_rejected(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
cache = _shared_cache()
resp = _build_signed_response(key_pem, cert_pem, email="bob@example.com")
first = await _acs(_b64(resp), cache)
assert first.email == "bob@example.com"
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), cache)
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_assertion_without_id_is_rejected(saml_env_idp_initiated):
"""An assertion with no ID attribute has no stable replay key. On the unsolicited
path there is no browser binding, so the consumed-assertion guard is the only replay
defense; a missing ID must be rejected rather than silently skipping the guard."""
class _AuthNoAssertionId:
def get_last_response_in_response_to(self):
return None
def get_last_response_xml(self):
return None
def get_last_assertion_id(self):
return None
with pytest.raises(HTTPException) as exc:
await SAMLAuthHandler._enforce_response_binding(
_AuthNoAssertionId(), _shared_cache(), None
)
assert exc.value.status_code == 401
assert "ID" in exc.value.detail
@pytest.mark.asyncio
async def test_unsolicited_response_rejected_when_disabled(saml_env, monkeypatch):
key_pem, cert_pem = saml_env
monkeypatch.setenv("SAML_ALLOW_UNSOLICITED", "false")
resp = _build_signed_response(key_pem, cert_pem)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
@pytest.mark.asyncio
async def test_invalid_email_in_assertion_is_rejected_cleanly(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(
key_pem,
cert_pem,
email="not-an-email",
attributes={"email": ["not-an-email"], "givenName": ["X"]},
)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), _shared_cache())
assert exc.value.status_code == 401
assert "invalid subject or email" in exc.value.detail
@pytest.mark.asyncio
async def test_email_less_assertion_rejected_when_domain_restriction_configured(
saml_env_idp_initiated, monkeypatch
):
key_pem, cert_pem = saml_env_idp_initiated
monkeypatch.setenv("ALLOWED_EMAIL_DOMAINS", "example.com")
resp = _build_signed_response(
key_pem,
cert_pem,
email="opaque-persistent-id-123",
attributes={"givenName": ["Alice"]},
)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), _shared_cache())
assert exc.value.status_code == 401
assert "ALLOWED_EMAIL_DOMAINS" in exc.value.detail
@pytest.mark.asyncio
async def test_email_less_assertion_allowed_without_domain_restriction(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(
key_pem,
cert_pem,
email="opaque-persistent-id-123",
attributes={"givenName": ["Alice"]},
)
result = await _acs(_b64(resp), _shared_cache())
assert result.email is None
assert result.id == "opaque-persistent-id-123"
@pytest.mark.asyncio
async def test_custom_email_attribute_override(saml_env_idp_initiated, monkeypatch):
key_pem, cert_pem = saml_env_idp_initiated
monkeypatch.setenv("SAML_ATTRIBUTE_EMAIL", "corpMail")
resp = _build_signed_response(
key_pem,
cert_pem,
email="ignored@example.com",
attributes={
"corpMail": ["real@corp.example.com"],
"givenName": ["Real"],
},
)
result = await _acs(_b64(resp), _shared_cache())
assert result.email == "real@corp.example.com"
@pytest.mark.asyncio
async def test_team_ids_extracted_from_groups_attribute(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(
key_pem,
cert_pem,
attributes={
"email": ["carol@example.com"],
"groups": ["team-a", "team-b"],
},
)
result = await _acs(_b64(resp), _shared_cache())
assert result.team_ids == ["team-a", "team-b"]
@pytest.mark.asyncio
async def test_build_login_redirect_targets_idp_and_caches_request_id(saml_env):
cache = DualCache()
redirect = await SAMLAuthHandler.build_login_redirect(_fake_request(), cache)
location = redirect.headers["location"]
assert location.startswith(SSO_URL)
assert "SAMLRequest=" in location
cached = [
k
for k in cache.in_memory_cache.cache_dict
if k.startswith(_SAML_AUTHN_REQUEST_CACHE_PREFIX)
]
assert len(cached) == 1
request_id = cached[0].split(":", 1)[1]
set_cookie = redirect.headers["set-cookie"]
assert f"{_SAML_AUTHN_STATE_COOKIE}={request_id}" in set_cookie
assert "httponly" in set_cookie.lower()
@pytest.mark.asyncio
async def test_sp_metadata_contains_acs_and_entity_id(saml_env):
metadata = await SAMLAuthHandler.build_sp_metadata(_fake_request(), DualCache())
assert ACS in metadata
assert SP_ENTITY in metadata
assert "AssertionConsumerService" in metadata
def test_replay_guard_ttl_tracks_assertion_validity():
class _Auth:
def __init__(self, not_on_or_after):
self._not_on_or_after = not_on_or_after
def get_last_assertion_not_on_or_after(self):
return self._not_on_or_after
now = int(time.time())
long_lived = SAMLAuthHandler._replay_guard_ttl(_Auth(now + 7200))
assert long_lived >= 7200
short_lived = SAMLAuthHandler._replay_guard_ttl(_Auth(now + 60))
assert short_lived == _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS
missing = SAMLAuthHandler._replay_guard_ttl(_Auth(None))
assert missing == _SAML_REPLAY_GUARD_DEFAULT_TTL_SECONDS
capped = SAMLAuthHandler._replay_guard_ttl(_Auth(now + 10 * 86400))
assert capped == _SAML_REPLAY_GUARD_MAX_TTL_SECONDS
def test_is_saml_configured_reflects_env(monkeypatch):
monkeypatch.delenv("SAML_IDP_METADATA_URL", raising=False)
monkeypatch.delenv("SAML_IDP_METADATA_XML", raising=False)
assert SAMLAuthHandler.is_saml_configured() is False
monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml")
assert SAMLAuthHandler.is_saml_configured() is True
@pytest.mark.asyncio
async def test_idp_initiated_rejected_without_shared_cache(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
resp = _build_signed_response(key_pem, cert_pem)
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), DualCache())
assert exc.value.status_code == 401
assert "shared Redis cache" in exc.value.detail
@pytest.mark.asyncio
async def test_idp_initiated_replay_rejected_across_workers(saml_env_idp_initiated):
key_pem, cert_pem = saml_env_idp_initiated
shared_store = InMemoryCache()
worker_one = _shared_cache(shared_store)
worker_two = _shared_cache(shared_store)
resp = _build_signed_response(key_pem, cert_pem, email="bob@example.com")
first = await _acs(_b64(resp), worker_one)
assert first.email == "bob@example.com"
with pytest.raises(HTTPException) as exc:
await _acs(_b64(resp), worker_two)
assert exc.value.status_code == 401
class _FakeChunkedRequest:
def __init__(self, chunks, content_length=None):
self._chunks = chunks
self.headers = {} if content_length is None else {"content-length": content_length}
async def stream(self):
for chunk in self._chunks:
yield chunk
@pytest.mark.asyncio
async def test_read_acs_post_data_parses_form():
body = b"SAMLResponse=abc123&RelayState=%2Fui%2F"
request = _FakeChunkedRequest([body], content_length=str(len(body)))
post_data = await SAMLAuthHandler.read_acs_post_data(cast(Request, request))
assert post_data == {"SAMLResponse": "abc123", "RelayState": "/ui/"}
@pytest.mark.asyncio
async def test_read_acs_post_data_rejects_oversized_content_length():
request = _FakeChunkedRequest([b""], content_length=str(_SAML_MAX_POST_BYTES + 1))
with pytest.raises(HTTPException) as exc:
await SAMLAuthHandler.read_acs_post_data(cast(Request, request))
assert exc.value.status_code == 413
@pytest.mark.asyncio
async def test_read_acs_post_data_rejects_oversized_stream_without_content_length():
chunk = b"a" * (1024 * 1024)
chunk_count = _SAML_MAX_POST_BYTES // len(chunk) + 2
request = _FakeChunkedRequest([chunk] * chunk_count)
with pytest.raises(HTTPException) as exc:
await SAMLAuthHandler.read_acs_post_data(cast(Request, request))
assert exc.value.status_code == 413

View file

@ -7391,6 +7391,71 @@ async def test_legacy_login_page_hides_credentials_hint_via_general_settings():
assert "MASTER_KEY" not in body
@pytest.mark.asyncio
async def test_saml_callback_blocked_when_admin_ui_disabled():
"""An IdP-initiated assertion must not mint a UI session when the admin UI is
disabled; the ACS enforces DISABLE_ADMIN_UI like the SP-initiated login route."""
from litellm.proxy.management_endpoints.ui_sso import saml_callback
with patch.dict(os.environ, {"DISABLE_ADMIN_UI": "true"}):
response = await saml_callback(SimpleNamespace(cookies={}))
assert response.status_code == 200
assert "Admin UI is Disabled" in response.body.decode()
@pytest.mark.asyncio
async def test_saml_callback_enforces_free_sso_user_limit_after_validation():
"""An IdP-initiated assertion must not bypass the >5 free-SSO-user Enterprise gate
that /sso/key/generate enforces; the ACS re-checks it after validating the assertion,
so the entitlement DB query never runs on unvalidated input."""
from litellm.proxy._types import ProxyException
from litellm.proxy.management_endpoints.ui_sso import saml_callback
from litellm.proxy.management_endpoints.types import CustomOpenID
call_order: list[str] = []
async def _fake_handle_acs(**kwargs):
call_order.append("validate")
return CustomOpenID(
id="dana@litellm.ai",
email="dana@litellm.ai",
first_name=None,
last_name=None,
display_name="dana",
picture=None,
provider="saml",
team_ids=[],
user_role=None,
)
async def _fake_count_billable_users():
call_order.append("count")
return 6
async def _stream():
yield b"SAMLResponse=signed-response"
request_double = SimpleNamespace(cookies={}, headers={}, stream=_stream)
with patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}), patch(
"litellm.proxy.proxy_server.premium_user", False
), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch(
"litellm.proxy.proxy_server.master_key", "sk-1234"
), patch(
"litellm.proxy.management_endpoints.sso.saml_sso.SAMLAuthHandler.handle_acs",
new=_fake_handle_acs,
), patch(
"litellm.repositories.user_repository.UserRepository.count_billable_users",
new=AsyncMock(side_effect=_fake_count_billable_users),
):
with pytest.raises(ProxyException) as exc:
await saml_callback(request_double)
assert str(exc.value.code) == "403"
assert call_order == ["validate", "count"]
@pytest.mark.asyncio
async def test_cli_poll_key_tolerates_missing_user_row():
"""The CLI poll must still mint the JWT when the user lookup raises,

View file

@ -1504,50 +1504,6 @@ def test_team_info_masking():
assert "public-test-key" not in str(exc_info.value)
def test_embedding_input_array_of_tokens(client_no_auth):
"""
Test to bypass decoding input as array of tokens for selected providers
Ref: https://github.com/BerriAI/litellm/issues/10113
"""
from litellm.proxy import proxy_server
# The client_no_auth fixture should initialize the router
# Assert this to catch any router initialization regressions
assert proxy_server.llm_router is not None, (
"llm_router is None after client_no_auth fixture initialized. "
"This indicates a router initialization issue that should be investigated."
)
try:
with mock.patch.object(
proxy_server.llm_router,
"aembedding",
return_value=example_embedding_result,
) as mock_aembedding:
test_data = {
"model": "vllm_embed_model",
"input": [[2046, 13269, 158208]],
}
response = client_no_auth.post("/v1/embeddings", json=test_data)
# Assert that aembedding was called, and that input was not modified
mock_aembedding.assert_called_once()
call_args, call_kwargs = mock_aembedding.call_args
assert call_kwargs["model"] == "vllm_embed_model"
assert call_kwargs["input"] == [[2046, 13269, 158208]]
assert response.status_code == 200
result = response.json()
print(len(result["data"][0]["embedding"]))
assert (
len(result["data"][0]["embedding"]) > 10
) # this usually has len==1536 so
except Exception as e:
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
@pytest.mark.asyncio
async def test_get_all_team_models():
"""

View file

@ -12,24 +12,25 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
@pytest.mark.parametrize(
"route_type",
"route_type, required_body_params",
[
"atext_completion",
"acompletion",
"aembedding",
"aimage_generation",
"aspeech",
"atranscription",
"amoderation",
"arerank",
("atext_completion", {}),
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
("aembedding", {"input": "Hello"}),
("aimage_generation", {}),
("aspeech", {}),
("atranscription", {}),
("amoderation", {}),
("arerank", {}),
],
)
@pytest.mark.asyncio
async def test_route_request_dynamic_credentials(route_type):
async def test_route_request_dynamic_credentials(route_type, required_body_params):
data = {
"model": "openai/gpt-4o-mini-2024-07-18",
"api_key": "my-bad-key",
"api_base": "https://api.openai.com/v1 ",
**required_body_params,
}
llm_router = MagicMock()
# Ensure that the dynamic method exists on the llm_router mock.
@ -887,3 +888,59 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
call_kwargs = llm_router.acompletion.call_args[1]
assert call_kwargs["enable_tag_filtering"] is True
@pytest.mark.parametrize(
"route_type, param, route",
[
("acompletion", "messages", "/chat/completions"),
("aembedding", "input", "/embeddings"),
],
)
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None}])
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
from litellm.proxy.route_llm_request import (
ProxyMissingRequiredParamError,
raise_if_required_body_param_missing,
)
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
raise_if_required_body_param_missing(route_type=route_type, data={"model": "gpt-4o", **data_extra})
assert exc_info.value.status_code == 400
assert exc_info.value.param == param
assert exc_info.value.type == "invalid_request_error"
assert exc_info.value.detail == {"error": f"{route}: Missing required parameter: '{param}'."}
@pytest.mark.parametrize(
"route_type, data",
[
("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}),
("acompletion", {"model": "gpt-4o", "messages": []}),
("atext_completion", {"model": "gpt-4o"}),
("aembedding", {"model": "text-embedding-3-small", "input": "hi"}),
("arerank", {"model": "rerank-model"}),
("aimage_generation", {"model": "dall-e-3"}),
],
)
def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data):
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
raise_if_required_body_param_missing(route_type=route_type, data=data)
@pytest.mark.asyncio
async def test_route_request_rejects_chat_completion_without_messages():
"""A /chat/completions body without `messages` used to splat into
Router.acompletion() and surface the resulting TypeError as a 500."""
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
llm_router = MagicMock()
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
await route_request({"model": "gpt-4o"}, llm_router, None, "acompletion")
assert exc_info.value.status_code == 400
assert exc_info.value.param == "messages"
llm_router.acompletion.assert_not_called()

View file

@ -617,6 +617,75 @@ class TestProxySettingEndpoints:
create_sso_settings = json.loads(create_data["sso_settings"])
assert create_sso_settings["google_client_id"] == "new_google_client_id"
def test_update_sso_settings_maps_saml_fields_to_env_vars(
self, mock_proxy_config, mock_auth, monkeypatch
):
"""SAML settings entered in the admin UI must be applied as the SAML_* env
vars the SAML handler reads, and the allow-unsolicited toggle must map to
the 'true'/'false' string the handler expects."""
import json
import os
from unittest.mock import AsyncMock, MagicMock
monkeypatch.setenv("LITELLM_SALT_KEY", "test_salt_key")
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
mock_prisma = MagicMock()
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_ssoconfig.upsert = AsyncMock()
mock_prisma.db.litellm_config = MagicMock()
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_config.update = AsyncMock()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
from litellm.proxy.proxy_server import proxy_config
monkeypatch.setattr(
proxy_config,
"_encrypt_env_variables",
lambda environment_variables: environment_variables,
)
for var in (
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
"SAML_SP_ENTITY_ID",
"SAML_ALLOW_UNSOLICITED",
):
monkeypatch.delenv(var, raising=False)
new_sso_settings = {
"saml_idp_metadata_url": "https://idp.example.com/metadata",
"saml_sp_entity_id": "https://proxy.example.com/sso/saml/metadata",
"saml_allow_unsolicited": "true",
"proxy_base_url": "https://proxy.example.com",
"user_email": "admin@example.com",
}
try:
response = client.patch("/update/sso_settings", json=new_sso_settings)
assert response.status_code == 200
assert os.environ.get("SAML_IDP_METADATA_URL") == "https://idp.example.com/metadata"
assert os.environ.get("SAML_SP_ENTITY_ID") == "https://proxy.example.com/sso/saml/metadata"
assert os.environ.get("SAML_ALLOW_UNSOLICITED") == "true"
assert "SAML_IDP_METADATA_XML" not in os.environ
stored = json.loads(
mock_prisma.db.litellm_ssoconfig.upsert.call_args.kwargs["data"]["create"]["sso_settings"]
)
assert stored["saml_idp_metadata_url"] == "https://idp.example.com/metadata"
assert stored["saml_allow_unsolicited"] == "true"
finally:
for var in (
"SAML_IDP_METADATA_URL",
"SAML_IDP_METADATA_XML",
"SAML_SP_ENTITY_ID",
"SAML_ALLOW_UNSOLICITED",
):
os.environ.pop(var, None)
def test_update_sso_settings_audits_when_env_cleanup_fails(
self, mock_proxy_config, mock_auth, monkeypatch
):

View file

@ -626,6 +626,7 @@ def _moderation_guardrail() -> MagicMock:
cb.should_run_guardrail = MagicMock(return_value=True)
cb.async_moderation_hook = AsyncMock(return_value=None)
cb.async_post_call_success_hook = AsyncMock(return_value=None)
cb.run_in_parallel = False
return cb

View file

@ -28,6 +28,7 @@ def _make_guardrail(name="g", should_run=True, override=None):
cb.event_hook = GuardrailEventHooks.post_call
cb.should_run_guardrail = MagicMock(return_value=should_run)
cb.async_post_call_success_hook = AsyncMock(return_value=override)
cb.run_in_parallel = False
return cb

View file

@ -198,6 +198,54 @@ async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error():
assert "not supported" in str(excinfo.value).lower()
@pytest.mark.asyncio
async def test_aresponses_drops_stream_options():
"""The Responses API rejects include_usage, so include_usage-only stream_options must never reach the wire."""
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(
_minimal_responses_api_payload("resp_stream_options_test", "gpt-5.5"), 200
)
await litellm.aresponses(
model="openai/gpt-5.5",
api_key="fake-api-key",
input="hi",
stream_options={"include_usage": True},
)
mock_post.assert_called_once()
post_kwargs = mock_post.call_args.kwargs
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
assert "stream_options" not in request_body
@pytest.mark.asyncio
async def test_aresponses_keeps_include_obfuscation_in_stream_options():
"""include_obfuscation is a valid Responses API stream option and must survive the include_usage strip."""
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(
_minimal_responses_api_payload("resp_stream_options_obfuscation", "gpt-5.5"), 200
)
await litellm.aresponses(
model="openai/gpt-5.5",
api_key="fake-api-key",
input="hi",
stream_options={"include_usage": True, "include_obfuscation": False},
)
mock_post.assert_called_once()
post_kwargs = mock_post.call_args.kwargs
request_body = post_kwargs["json"] if "json" in post_kwargs else json.loads(post_kwargs["data"])
assert request_body["stream_options"] == {"include_obfuscation": False}
@pytest.mark.asyncio
async def test_aresponses_request_level_drop_params_drops_bedrock_mantle_service_tier(
monkeypatch,

View file

@ -1088,11 +1088,6 @@
"count": 1
}
},
"src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/playground/components/chat_ui/A2AMetrics.tsx": {
"no-restricted-imports": {
"count": 1

View file

@ -24,6 +24,10 @@ export interface SSOSettingsValues {
generic_authorization_endpoint: string | null;
generic_token_endpoint: string | null;
generic_userinfo_endpoint: string | null;
saml_idp_metadata_url: string | null;
saml_idp_metadata_xml: string | null;
saml_sp_entity_id: string | null;
saml_allow_unsolicited: string | null;
generic_scope: string | null;
proxy_base_url: string | null;
user_email: string | null;

View file

@ -1,19 +1,14 @@
import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels";
import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters";
import { InfoCircleOutlined } from "@ant-design/icons";
import { Form, Input, Modal, Select as Select2, Tooltip } from "antd";
import { useQueryClient } from "@tanstack/react-query";
import React, { useState } from "react";
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { organizationCreateCall, organizationDeleteCall } from "@/components/networking";
import { organizationDeleteCall } from "@/components/networking";
import { OrgCreateDialog } from "@/components/organization/org-create/OrgCreateDialog";
import OrganizationInfoView from "@/components/organization/organization_view";
import NumericalInput from "@/components/shared/numerical_input";
import { Button } from "@/components/ui/button";
import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector";
import OrganizationsTable from "./OrganizationsTable";
@ -30,7 +25,6 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
const [orgToDelete, setOrgToDelete] = useState<string | null>(null);
const [isDeleting, setIsDeleting] = useState(false);
const [isOrgModalVisible, setIsOrgModalVisible] = useState(false);
const [form] = Form.useForm();
const [showFilters, setShowFilters] = useState(false);
const [filters, setFilters] = useState<FilterState>({ org_id: "", org_alias: "" });
@ -83,48 +77,6 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
setOrgToDelete(null);
};
const handleCreate = async (values: any) => {
try {
if (!accessToken) return;
// Transform allowed_vector_store_ids and allowed_mcp_servers_and_groups into object_permission
if (
(values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) ||
(values.allowed_mcp_servers_and_groups &&
(values.allowed_mcp_servers_and_groups.servers?.length > 0 ||
values.allowed_mcp_servers_and_groups.accessGroups?.length > 0))
) {
values.object_permission = {};
if (values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) {
values.object_permission.vector_stores = values.allowed_vector_store_ids;
delete values.allowed_vector_store_ids;
}
if (values.allowed_mcp_servers_and_groups) {
if (values.allowed_mcp_servers_and_groups.servers?.length > 0) {
values.object_permission.mcp_servers = values.allowed_mcp_servers_and_groups.servers;
}
if (values.allowed_mcp_servers_and_groups.accessGroups?.length > 0) {
values.object_permission.mcp_access_groups = values.allowed_mcp_servers_and_groups.accessGroups;
}
delete values.allowed_mcp_servers_and_groups;
}
}
await organizationCreateCall(accessToken, values);
NotificationsManager.success("Organization created successfully");
setIsOrgModalVisible(false);
form.resetFields();
await refetchOrganizations();
} catch (error) {
console.error("Error creating organization:", error);
}
};
const handleCancel = () => {
setIsOrgModalVisible(false);
form.resetFields();
};
if (!premiumUser) {
return (
<div className="mx-4 mt-4">
@ -190,97 +142,7 @@ const OrganizationsPanel: React.FC<OrganizationsPanelProps> = ({ userRole, acces
</>
)}
<Modal title="Create Organization" visible={isOrgModalVisible} width={800} footer={null} onCancel={handleCancel}>
<Form form={form} onFinish={handleCreate} labelCol={{ span: 8 }} wrapperCol={{ span: 16 }} labelAlign="left">
<Form.Item
label="Organization Name"
name="organization_alias"
rules={[
{
required: true,
message: "Please input an organization name",
},
]}
>
<Input placeholder="" />
</Form.Item>
<Form.Item label="Models" name="models">
<ModelSelect
options={{ showAllProxyModelsOverride: true, includeSpecialOptions: true }}
value={form.getFieldValue("models")}
onChange={(values) => form.setFieldValue("models", values)}
context="organization"
/>
</Form.Item>
<Form.Item label="Max Budget (USD)" name="max_budget">
<NumericalInput step={0.01} precision={2} width={200} />
</Form.Item>
<Form.Item label="Reset Budget" name="budget_duration">
<Select2 defaultValue={null} placeholder="n/a">
<Select2.Option value="24h">daily</Select2.Option>
<Select2.Option value="7d">weekly</Select2.Option>
<Select2.Option value="30d">monthly</Select2.Option>
</Select2>
</Form.Item>
<Form.Item label="Tokens per minute Limit (TPM)" name="tpm_limit">
<NumericalInput step={1} width={400} />
</Form.Item>
<Form.Item label="Requests per minute Limit (RPM)" name="rpm_limit">
<NumericalInput step={1} width={400} />
</Form.Item>
<Form.Item
label={
<span>
Allowed Vector Stores{" "}
<Tooltip title="Select which vector stores this organization can access by default. Leave empty for access to all vector stores">
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</Tooltip>
</span>
}
name="allowed_vector_store_ids"
className="mt-4"
help="Select vector stores this organization can access. Leave empty for access to all vector stores"
>
<VectorStoreSelector
onChange={(values) => form.setFieldValue("allowed_vector_store_ids", values)}
value={form.getFieldValue("allowed_vector_store_ids")}
accessToken={accessToken || ""}
placeholder="Select vector stores (optional)"
/>
</Form.Item>
<Form.Item
label={
<span>
Allowed MCP Servers{" "}
<Tooltip title="Select which MCP servers and access groups this organization can access by default.">
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</Tooltip>
</span>
}
name="allowed_mcp_servers_and_groups"
className="mt-4"
help="Select MCP servers and access groups this organization can access."
>
<MCPServerSelector
onChange={(values) => form.setFieldValue("allowed_mcp_servers_and_groups", values)}
value={form.getFieldValue("allowed_mcp_servers_and_groups")}
accessToken={accessToken || ""}
placeholder="Select MCP servers and access groups (optional)"
/>
</Form.Item>
<Form.Item label="Metadata" name="metadata">
<Input.TextArea rows={4} />
</Form.Item>
<div style={{ textAlign: "right", marginTop: "10px" }}>
<Button type="submit">Create Organization</Button>
</div>
</Form>
</Modal>
<OrgCreateDialog open={isOrgModalVisible} onOpenChange={setIsOrgModalVisible} accessToken={accessToken || ""} />
<DeleteResourceModal
isOpen={isDeleteModalOpen}

View file

@ -25,8 +25,17 @@ vi.mock("@/components/networking", () => ({
// Mock the child components to simplify testing
vi.mock("@/components/activity_metrics", () => ({
ActivityMetrics: () => <div>Activity Metrics</div>,
processActivityData: () => ({ data: [], metadata: {} }),
ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => (
<div>
<span>Activity Metrics</span>
<span>{`metrics-source:${modelMetrics?.__source ?? "none"}`}</span>
</div>
),
processActivityData: (_data: unknown, key: string) => ({ __source: key }),
}));
vi.mock("../EndpointUsage/EndpointUsage", () => ({
default: () => <div>Endpoint Usage Panel</div>,
}));
vi.mock("@/components/UsagePage/components/EntityUsage/TopKeyView", () => ({
@ -481,6 +490,54 @@ describe("EntityUsage", () => {
expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument();
});
const selectedPanels = (container: HTMLElement) =>
Array.from(container.querySelectorAll("div.tremor-TabPanel-root")).filter(
(panel) => panel.getAttribute("aria-selected") === "true",
);
it.each([
["Cost", "Tag Spend Overview"],
["Model Activity", "metrics-source:models"],
["Key Activity", "metrics-source:api_keys"],
["Endpoint Activity", "Endpoint Usage Panel"],
])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => {
const { container } = render(<EntityUsage {...defaultProps} />);
await waitFor(() => {
expect(mockTagDailyActivityCall).toHaveBeenCalled();
});
act(() => {
fireEvent.click(screen.getByText(tabLabel));
});
const selected = selectedPanels(container);
expect(selected).toHaveLength(1);
expect(selected[0].textContent).toContain(marker);
});
it.each([
["Cost", "Team Spend Overview"],
["Model Activity", "metrics-source:models"],
["Agent Activity", "metrics-source:entities"],
["Key Activity", "metrics-source:api_keys"],
["Endpoint Activity", "Endpoint Usage Panel"],
])("shows only the %s panel for the team entity type", async (tabLabel, marker) => {
const { container } = render(<EntityUsage {...defaultProps} entityType="team" />);
await waitFor(() => {
expect(mockTeamDailyActivityCall).toHaveBeenCalled();
});
act(() => {
fireEvent.click(screen.getByText(tabLabel));
});
const selected = selectedPanels(container);
expect(selected).toHaveLength(1);
expect(selected[0].textContent).toContain(marker);
});
it("should handle empty data gracefully", async () => {
const emptyData = {
results: [],

View file

@ -25,7 +25,7 @@ import {
} from "@tremor/react";
import { ExportOutlined, LoadingOutlined } from "@ant-design/icons";
import { Alert, Button } from "antd";
import React, { useMemo, useState } from "react";
import React, { type ReactNode, useMemo, useState } from "react";
import TeamMultiSelect from "@/components/common_components/team_multi_select";
import { ActivityMetrics, processActivityData } from "@/components/activity_metrics";
import { UsageExportHeader } from "@/components/EntityUsageExport";
@ -406,6 +406,304 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1);
const costPanel = (
<Grid numItems={2} className="gap-2 w-full">
{/* Total Spend Card */}
<Col numColSpan={2}>
<Card>
<Title>{capitalizedEntityLabel} Spend Overview</Title>
<Grid numItems={5} className="gap-4 mt-4">
<Card>
<Title>Total Spend</Title>
<Text className="text-2xl font-bold mt-2">
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
</Text>
</Card>
<Card>
<Title>Total Requests</Title>
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_api_requests.toLocaleString()}</Text>
</Card>
<Card>
<Title>Successful Requests</Title>
<Text className="text-2xl font-bold mt-2 text-green-600">
{spendData.metadata.total_successful_requests.toLocaleString()}
</Text>
</Card>
<Card>
<Title>Failed Requests</Title>
<Text className="text-2xl font-bold mt-2 text-red-600">
{spendData.metadata.total_failed_requests.toLocaleString()}
</Text>
</Card>
<Card>
<Title>Total Tokens</Title>
<Text className="text-2xl font-bold mt-2">{spendData.metadata.total_tokens.toLocaleString()}</Text>
</Card>
</Grid>
</Card>
</Col>
{/* Daily Spend Chart */}
<Col numColSpan={2}>
<ShadcnCard>
<CardHeader>
<CardTitle className="text-base font-semibold">Daily Spend</CardTitle>
</CardHeader>
<CardContent>
<BarChart
data={[...spendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime())}
index="date"
categories={["metrics.spend"]}
colors={["cyan"]}
valueFormatter={valueFormatterSpend}
yAxisWidth={100}
showLegend={false}
customTooltip={({ payload, active }) => {
if (!active || !payload?.[0]) return null;
const data = payload[0].payload;
const entityCount = Object.keys(data.breakdown.entities || {}).length;
return (
<div className="bg-white p-4 shadow-lg rounded-lg border">
<p className="font-bold">{data.date}</p>
<p className="text-cyan-500">Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}</p>
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
<p className="text-gray-600">Total Tokens: {data.metrics.total_tokens}</p>
<p className="text-gray-600">
Total {capitalizedEntityLabel}s: {entityCount}
</p>
<div className="mt-2 border-t pt-2">
<p className="font-semibold">Spend by {capitalizedEntityLabel}:</p>
{Object.entries(data.breakdown.entities || {})
.sort(([, a], [, b]) => {
const spendA = (a as EntityMetrics).metrics.spend;
const spendB = (b as EntityMetrics).metrics.spend;
return spendB - spendA;
})
.slice(0, 5)
.map(([entity, entityData]) => {
const metrics = entityData as EntityMetrics;
return (
<p key={entity} className="text-sm text-gray-600">
{getEntityLabel(entity, metrics.metadata)}: $
{formatNumberWithCommas(metrics.metrics.spend, 2)}
</p>
);
})}
{entityCount > 5 && <p className="text-sm text-gray-500 italic">...and {entityCount - 5} more</p>}
</div>
</div>
);
}}
/>
</CardContent>
</ShadcnCard>
</Col>
{/* Entity Breakdown Section */}
<Col numColSpan={2}>
<Card>
<div className="flex flex-col space-y-4">
<div className="flex flex-col space-y-2">
<Title>Spend Per {capitalizedEntityLabel}</Title>
<Subtitle className="text-xs">Showing Top 5 by Spend</Subtitle>
<div className="flex items-center text-sm text-gray-500">
<span>Get Started by Tracking cost per {capitalizedEntityLabel} </span>
<a
href="https://docs.litellm.ai/docs/proxy/enterprise#spend-tracking"
className="text-blue-500 hover:text-blue-700 ml-1"
>
here
</a>
</div>
</div>
<Grid numItems={2} className="gap-6">
<Col numColSpan={1}>
<BarChart
className="mt-4 h-52"
data={getProcessedEntityBreakdownForChart()}
index="metadata.alias_display"
categories={["metrics.spend"]}
colors={["cyan"]}
valueFormatter={valueFormatterSpend}
layout="vertical"
showLegend={false}
yAxisWidth={150}
customTooltip={({ payload, active }) => {
if (!active || !payload?.[0]) return null;
const data = payload[0].payload;
return (
<div className="bg-white p-4 shadow-lg rounded-lg border">
<p className="font-bold">{data.metadata.alias}</p>
<p className="text-cyan-500">Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}</p>
<p className="text-gray-600">Requests: {data.metrics.api_requests.toLocaleString()}</p>
<p className="text-green-600">
Successful: {data.metrics.successful_requests.toLocaleString()}
</p>
<p className="text-red-600">Failed: {data.metrics.failed_requests.toLocaleString()}</p>
<p className="text-gray-600">Tokens: {data.metrics.total_tokens.toLocaleString()}</p>
</div>
);
}}
/>
</Col>
<Col numColSpan={1}>
<div className="h-52 overflow-y-auto">
<Table>
<TableHead>
<TableRow>
<TableHeaderCell>{capitalizedEntityLabel}</TableHeaderCell>
<TableHeaderCell>Spend</TableHeaderCell>
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
<TableHeaderCell>Tokens</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{getEntityBreakdown()
.filter((entity) => entity.metrics.spend > 0)
.map((entity) => (
<TableRow key={entity.metadata.id}>
<TableCell>{entity.metadata.alias}</TableCell>
<TableCell>
<MoneyCell value={entity.metrics.spend} decimals={4} />
</TableCell>
<TableCell className="text-green-600">
{entity.metrics.successful_requests.toLocaleString()}
</TableCell>
<TableCell className="text-red-600">
{entity.metrics.failed_requests.toLocaleString()}
</TableCell>
<TableCell>{entity.metrics.total_tokens.toLocaleString()}</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
</Col>
</Grid>
</div>
</Card>
</Col>
{/* Top API Keys */}
<Col numColSpan={1}>
<Card>
<Title>Top Virtual Keys</Title>
<TopKeyView
topKeys={getTopAPIKeys()}
teams={null}
showTags={entityType === "tag"}
topKeysLimit={topKeysLimit}
setTopKeysLimit={setTopKeysLimit}
/>
</Card>
</Col>
{/* Top Models */}
<Col numColSpan={1}>
<Card>
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
<TopModelView
topModels={getTopModels()}
topModelsLimit={topModelsLimit}
setTopModelsLimit={setTopModelsLimit}
/>
</Card>
</Col>
{/* Top Agents - only for team entity type */}
{entityType === "team" && (
<Col numColSpan={2}>
<Card>
<Title>Top Agents Driving Spend</Title>
<TopModelView
topModels={getTopAgents()}
topModelsLimit={topAgentsLimit}
setTopModelsLimit={setTopAgentsLimit}
/>
</Card>
</Col>
)}
{/* Spend by Provider */}
<Col numColSpan={2}>
<Card>
<div className="flex flex-col space-y-4">
<Title>Provider Usage</Title>
<Grid numItems={2}>
<Col numColSpan={1}>
<DonutChart
className="mt-4 h-40"
data={getProviderSpend()}
index="provider"
category="spend"
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
colors={["cyan", "blue", "indigo", "violet", "purple"]}
showLabel
startAngle={90}
endAngle={-270}
/>
</Col>
<Col numColSpan={1}>
<Table>
<TableHead>
<TableRow>
<TableHeaderCell>Provider</TableHeaderCell>
<TableHeaderCell>Spend</TableHeaderCell>
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
<TableHeaderCell>Tokens</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{getProviderSpend().map((provider) => (
<TableRow key={provider.provider}>
<TableCell>
<div className="flex items-center space-x-2">
{provider.provider && <Logo provider={provider.provider} className="w-4 h-4" />}
<span>{provider.provider}</span>
</div>
</TableCell>
<TableCell>
<MoneyCell value={provider.spend} decimals={2} />
</TableCell>
<TableCell className="text-green-600">
{provider.successful_requests.toLocaleString()}
</TableCell>
<TableCell className="text-red-600">{provider.failed_requests.toLocaleString()}</TableCell>
<TableCell>{provider.tokens.toLocaleString()}</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</Col>
</Grid>
</div>
</Card>
</Col>
</Grid>
);
const tabs: readonly { key: string; label: string; content: ReactNode }[] = [
{ key: "cost", label: "Cost", content: costPanel },
{
key: "models",
label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity",
content: <ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
},
...(entityType === "team"
? [{ key: "agents", label: "Agent Activity", content: <ActivityMetrics modelMetrics={agentMetrics} /> }]
: []),
{
key: "keys",
label: "Key Activity",
content: <ActivityMetrics modelMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />,
},
{ key: "endpoints", label: "Endpoint Activity", content: <EndpointUsage userSpendData={spendData} /> },
];
return (
<div style={{ width: "100%" }} className="relative">
{isFetchingMore && (
@ -501,320 +799,14 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
/>
<TabGroup>
<TabList variant="solid" className="mt-1">
<Tab>Cost</Tab>
<Tab>{entityType === "agent" ? "Request / Token Consumption" : "Model Activity"}</Tab>
{entityType === "team" ? <Tab>Agent Activity</Tab> : <></>}
<Tab>Key Activity</Tab>
<Tab>Endpoint Activity</Tab>
{tabs.map(({ key, label }) => (
<Tab key={key}>{label}</Tab>
))}
</TabList>
<TabPanels>
<TabPanel>
<Grid numItems={2} className="gap-2 w-full">
{/* Total Spend Card */}
<Col numColSpan={2}>
<Card>
<Title>{capitalizedEntityLabel} Spend Overview</Title>
<Grid numItems={5} className="gap-4 mt-4">
<Card>
<Title>Total Spend</Title>
<Text className="text-2xl font-bold mt-2">
${formatNumberWithCommas(spendData.metadata.total_spend, 2)}
</Text>
</Card>
<Card>
<Title>Total Requests</Title>
<Text className="text-2xl font-bold mt-2">
{spendData.metadata.total_api_requests.toLocaleString()}
</Text>
</Card>
<Card>
<Title>Successful Requests</Title>
<Text className="text-2xl font-bold mt-2 text-green-600">
{spendData.metadata.total_successful_requests.toLocaleString()}
</Text>
</Card>
<Card>
<Title>Failed Requests</Title>
<Text className="text-2xl font-bold mt-2 text-red-600">
{spendData.metadata.total_failed_requests.toLocaleString()}
</Text>
</Card>
<Card>
<Title>Total Tokens</Title>
<Text className="text-2xl font-bold mt-2">
{spendData.metadata.total_tokens.toLocaleString()}
</Text>
</Card>
</Grid>
</Card>
</Col>
{/* Daily Spend Chart */}
<Col numColSpan={2}>
<ShadcnCard>
<CardHeader>
<CardTitle className="text-base font-semibold">Daily Spend</CardTitle>
</CardHeader>
<CardContent>
<BarChart
data={[...spendData.results].sort(
(a, b) => new Date(a.date).getTime() - new Date(b.date).getTime(),
)}
index="date"
categories={["metrics.spend"]}
colors={["cyan"]}
valueFormatter={valueFormatterSpend}
yAxisWidth={100}
showLegend={false}
customTooltip={({ payload, active }) => {
if (!active || !payload?.[0]) return null;
const data = payload[0].payload;
const entityCount = Object.keys(data.breakdown.entities || {}).length;
return (
<div className="bg-white p-4 shadow-lg rounded-lg border">
<p className="font-bold">{data.date}</p>
<p className="text-cyan-500">
Total Spend: ${formatNumberWithCommas(data.metrics.spend, 2)}
</p>
<p className="text-gray-600">Total Requests: {data.metrics.api_requests}</p>
<p className="text-gray-600">Successful: {data.metrics.successful_requests}</p>
<p className="text-gray-600">Failed: {data.metrics.failed_requests}</p>
<p className="text-gray-600">Total Tokens: {data.metrics.total_tokens}</p>
<p className="text-gray-600">
Total {capitalizedEntityLabel}s: {entityCount}
</p>
<div className="mt-2 border-t pt-2">
<p className="font-semibold">Spend by {capitalizedEntityLabel}:</p>
{Object.entries(data.breakdown.entities || {})
.sort(([, a], [, b]) => {
const spendA = (a as EntityMetrics).metrics.spend;
const spendB = (b as EntityMetrics).metrics.spend;
return spendB - spendA;
})
.slice(0, 5)
.map(([entity, entityData]) => {
const metrics = entityData as EntityMetrics;
return (
<p key={entity} className="text-sm text-gray-600">
{getEntityLabel(entity, metrics.metadata)}: $
{formatNumberWithCommas(metrics.metrics.spend, 2)}
</p>
);
})}
{entityCount > 5 && (
<p className="text-sm text-gray-500 italic">...and {entityCount - 5} more</p>
)}
</div>
</div>
);
}}
/>
</CardContent>
</ShadcnCard>
</Col>
{/* Entity Breakdown Section */}
<Col numColSpan={2}>
<Card>
<div className="flex flex-col space-y-4">
<div className="flex flex-col space-y-2">
<Title>Spend Per {capitalizedEntityLabel}</Title>
<Subtitle className="text-xs">Showing Top 5 by Spend</Subtitle>
<div className="flex items-center text-sm text-gray-500">
<span>Get Started by Tracking cost per {capitalizedEntityLabel} </span>
<a
href="https://docs.litellm.ai/docs/proxy/enterprise#spend-tracking"
className="text-blue-500 hover:text-blue-700 ml-1"
>
here
</a>
</div>
</div>
<Grid numItems={2} className="gap-6">
<Col numColSpan={1}>
<BarChart
className="mt-4 h-52"
data={getProcessedEntityBreakdownForChart()}
index="metadata.alias_display"
categories={["metrics.spend"]}
colors={["cyan"]}
valueFormatter={valueFormatterSpend}
layout="vertical"
showLegend={false}
yAxisWidth={150}
customTooltip={({ payload, active }) => {
if (!active || !payload?.[0]) return null;
const data = payload[0].payload;
return (
<div className="bg-white p-4 shadow-lg rounded-lg border">
<p className="font-bold">{data.metadata.alias}</p>
<p className="text-cyan-500">Spend: ${formatNumberWithCommas(data.metrics.spend, 4)}</p>
<p className="text-gray-600">Requests: {data.metrics.api_requests.toLocaleString()}</p>
<p className="text-green-600">
Successful: {data.metrics.successful_requests.toLocaleString()}
</p>
<p className="text-red-600">Failed: {data.metrics.failed_requests.toLocaleString()}</p>
<p className="text-gray-600">Tokens: {data.metrics.total_tokens.toLocaleString()}</p>
</div>
);
}}
/>
</Col>
<Col numColSpan={1}>
<div className="h-52 overflow-y-auto">
<Table>
<TableHead>
<TableRow>
<TableHeaderCell>{capitalizedEntityLabel}</TableHeaderCell>
<TableHeaderCell>Spend</TableHeaderCell>
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
<TableHeaderCell>Tokens</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{getEntityBreakdown()
.filter((entity) => entity.metrics.spend > 0)
.map((entity) => (
<TableRow key={entity.metadata.id}>
<TableCell>{entity.metadata.alias}</TableCell>
<TableCell>
<MoneyCell value={entity.metrics.spend} decimals={4} />
</TableCell>
<TableCell className="text-green-600">
{entity.metrics.successful_requests.toLocaleString()}
</TableCell>
<TableCell className="text-red-600">
{entity.metrics.failed_requests.toLocaleString()}
</TableCell>
<TableCell>{entity.metrics.total_tokens.toLocaleString()}</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
</Col>
</Grid>
</div>
</Card>
</Col>
{/* Top API Keys */}
<Col numColSpan={1}>
<Card>
<Title>Top Virtual Keys</Title>
<TopKeyView
topKeys={getTopAPIKeys()}
teams={null}
showTags={entityType === "tag"}
topKeysLimit={topKeysLimit}
setTopKeysLimit={setTopKeysLimit}
/>
</Card>
</Col>
{/* Top Models */}
<Col numColSpan={1}>
<Card>
<Title>{entityType === "agent" ? "Top Agents" : "Top Models"}</Title>
<TopModelView
topModels={getTopModels()}
topModelsLimit={topModelsLimit}
setTopModelsLimit={setTopModelsLimit}
/>
</Card>
</Col>
{/* Top Agents - only for team entity type */}
{entityType === "team" && (
<Col numColSpan={2}>
<Card>
<Title>Top Agents Driving Spend</Title>
<TopModelView
topModels={getTopAgents()}
topModelsLimit={topAgentsLimit}
setTopModelsLimit={setTopAgentsLimit}
/>
</Card>
</Col>
)}
{/* Spend by Provider */}
<Col numColSpan={2}>
<Card>
<div className="flex flex-col space-y-4">
<Title>Provider Usage</Title>
<Grid numItems={2}>
<Col numColSpan={1}>
<DonutChart
className="mt-4 h-40"
data={getProviderSpend()}
index="provider"
category="spend"
valueFormatter={(value) => `$${formatNumberWithCommas(value, 2)}`}
colors={["cyan", "blue", "indigo", "violet", "purple"]}
showLabel
startAngle={90}
endAngle={-270}
/>
</Col>
<Col numColSpan={1}>
<Table>
<TableHead>
<TableRow>
<TableHeaderCell>Provider</TableHeaderCell>
<TableHeaderCell>Spend</TableHeaderCell>
<TableHeaderCell className="text-green-600">Successful</TableHeaderCell>
<TableHeaderCell className="text-red-600">Failed</TableHeaderCell>
<TableHeaderCell>Tokens</TableHeaderCell>
</TableRow>
</TableHead>
<TableBody>
{getProviderSpend().map((provider) => (
<TableRow key={provider.provider}>
<TableCell>
<div className="flex items-center space-x-2">
{provider.provider && <Logo provider={provider.provider} className="w-4 h-4" />}
<span>{provider.provider}</span>
</div>
</TableCell>
<TableCell>
<MoneyCell value={provider.spend} decimals={2} />
</TableCell>
<TableCell className="text-green-600">
{provider.successful_requests.toLocaleString()}
</TableCell>
<TableCell className="text-red-600">
{provider.failed_requests.toLocaleString()}
</TableCell>
<TableCell>{provider.tokens.toLocaleString()}</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</Col>
</Grid>
</div>
</Card>
</Col>
</Grid>
</TabPanel>
<TabPanel>
<ActivityMetrics modelMetrics={modelMetrics} hidePromptCachingMetrics={entityType === "agent"} />
</TabPanel>
{entityType === "team" ? (
<TabPanel>
<ActivityMetrics modelMetrics={agentMetrics} />
</TabPanel>
) : (
<></>
)}
<TabPanel>
<ActivityMetrics modelMetrics={keyMetrics} hidePromptCachingMetrics={entityType === "agent"} />
</TabPanel>
<TabPanel>
<EndpointUsage userSpendData={spendData} />
</TabPanel>
{tabs.map(({ key, content }) => (
<TabPanel key={key}>{content}</TabPanel>
))}
</TabPanels>
</TabGroup>
</div>

View file

@ -1,5 +1,5 @@
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import { Form } from "antd";
import { Form, type FormInstance } from "antd";
import { describe, expect, it, vi } from "vitest";
import SSOModals from "./SSOModals";
@ -412,6 +412,76 @@ describe("SSOModals", () => {
expect(mockHandleShowInstructions).toHaveBeenCalled();
});
it("should submit SAML settings with the unsolicited toggle mapped to a 'true'/'false' string", async () => {
const mockHandleShowInstructions = vi.fn();
vi.mocked(updateSSOSettings).mockResolvedValue({});
vi.mocked(getSSOSettings).mockResolvedValue({ values: {} });
let formInstance: FormInstance | null = null;
const TestWrapper = () => {
const [form] = Form.useForm();
formInstance = form;
return (
<SSOModals
isAddSSOModalVisible={true}
isInstructionsModalVisible={false}
handleAddSSOOk={() => {}}
handleAddSSOCancel={() => {}}
handleShowInstructions={mockHandleShowInstructions}
handleInstructionsOk={() => {}}
handleInstructionsCancel={() => {}}
form={form}
accessToken="test-token"
ssoConfigured={false}
/>
);
};
render(<TestWrapper />);
await waitFor(() => {
expect(getSSOSettings).toHaveBeenCalledWith("test-token");
});
formInstance?.setFieldsValue({ sso_provider: "saml" });
await waitFor(() => {
expect(screen.getByLabelText("IdP Metadata URL")).toBeInTheDocument();
});
fireEvent.change(screen.getByLabelText("Proxy Admin Email"), {
target: { value: "admin@example.com" },
});
fireEvent.change(screen.getByLabelText("Proxy Base URL"), {
target: { value: "https://proxy.example.com" },
});
fireEvent.change(screen.getByLabelText("IdP Metadata URL"), {
target: { value: "https://idp.example.com/metadata" },
});
fireEvent.change(screen.getByLabelText("SP Entity ID"), {
target: { value: "https://proxy.example.com/sso/saml/metadata" },
});
fireEvent.click(screen.getByLabelText("Allow IdP-initiated (unsolicited) responses"));
fireEvent.click(screen.getByText("Save"));
await waitFor(() => {
expect(updateSSOSettings).toHaveBeenCalledWith(
"test-token",
expect.objectContaining({
sso_provider: "saml",
saml_idp_metadata_url: "https://idp.example.com/metadata",
saml_sp_entity_id: "https://proxy.example.com/sso/saml/metadata",
saml_allow_unsolicited: "true",
}),
);
});
expect(mockHandleShowInstructions).toHaveBeenCalled();
});
it("should show Clear button and clear SSO settings when configured", async () => {
const mockHandleAddSSOOk = vi.fn();
(updateSSOSettings as any).mockResolvedValue({});
@ -462,6 +532,10 @@ describe("SSOModals", () => {
generic_authorization_endpoint: null,
generic_token_endpoint: null,
generic_userinfo_endpoint: null,
saml_idp_metadata_url: null,
saml_idp_metadata_xml: null,
saml_sp_entity_id: null,
saml_allow_unsolicited: null,
generic_scope: null,
proxy_base_url: null,
user_email: null,

View file

@ -21,6 +21,17 @@ interface SSOModalsProps {
ssoConfigured?: boolean; // Add optional prop to indicate if SSO is configured
}
const detectSSOProvider = (values: Record<string, unknown>): string | null => {
if (values.google_client_id) return "google";
if (values.microsoft_client_id) return "microsoft";
if (values.generic_client_id) {
const authEndpoint =
typeof values.generic_authorization_endpoint === "string" ? values.generic_authorization_endpoint : "";
return authEndpoint.includes("okta") || authEndpoint.includes("auth0") ? "okta" : "generic";
}
if (values.saml_idp_metadata_url || values.saml_idp_metadata_xml) return "saml";
return null;
};
const SSOModals: React.FC<SSOModalsProps> = ({
isAddSSOModalVisible,
isInstructionsModalVisible,
@ -43,22 +54,7 @@ const SSOModals: React.FC<SSOModalsProps> = ({
const ssoData = await getSSOSettings(accessToken);
if (ssoData && ssoData.values) {
// Determine which SSO provider is configured
let selectedProvider = null;
if (ssoData.values.google_client_id) {
selectedProvider = "google";
} else if (ssoData.values.microsoft_client_id) {
selectedProvider = "microsoft";
} else if (ssoData.values.generic_client_id) {
// Check if it looks like Okta based on endpoints
if (
ssoData.values.generic_authorization_endpoint?.includes("okta") ||
ssoData.values.generic_authorization_endpoint?.includes("auth0")
) {
selectedProvider = "okta";
} else {
selectedProvider = "generic";
}
}
const selectedProvider = detectSSOProvider(ssoData.values);
// Extract role mappings if they exist
let roleMappingFields = {};
@ -89,6 +85,7 @@ const SSOModals: React.FC<SSOModalsProps> = ({
user_email: ssoData.values.user_email,
...ssoData.values,
...roleMappingFields,
saml_allow_unsolicited: ssoData.values.saml_allow_unsolicited === "true",
};
// Clear form first, then set values with a small delay to ensure proper initialization
@ -129,6 +126,10 @@ const SSOModals: React.FC<SSOModalsProps> = ({
...rest,
};
if (typeof payload.saml_allow_unsolicited === "boolean") {
payload.saml_allow_unsolicited = payload.saml_allow_unsolicited ? "true" : "false";
}
// Add role mappings if use_role_mappings is checked
if (use_role_mappings) {
// Helper function to split comma-separated string into array
@ -191,6 +192,10 @@ const SSOModals: React.FC<SSOModalsProps> = ({
generic_authorization_endpoint: null,
generic_token_endpoint: null,
generic_userinfo_endpoint: null,
saml_idp_metadata_url: null,
saml_idp_metadata_xml: null,
saml_sp_entity_id: null,
saml_allow_unsolicited: null,
generic_scope: null,
proxy_base_url: null,
user_email: null,

View file

@ -19,6 +19,7 @@ export interface SSOProviderConfig {
name: string;
placeholder?: string;
required?: boolean;
type?: "password" | "textarea" | "checkbox";
}>;
}
@ -90,6 +91,41 @@ export const ssoProviderConfigs: Record<string, SSOProviderConfig> = {
{ label: "Scopes", name: "generic_scope", placeholder: "openid email profile", required: false },
],
},
saml: {
envVarMap: {
saml_idp_metadata_url: "SAML_IDP_METADATA_URL",
saml_idp_metadata_xml: "SAML_IDP_METADATA_XML",
saml_sp_entity_id: "SAML_SP_ENTITY_ID",
saml_allow_unsolicited: "SAML_ALLOW_UNSOLICITED",
},
fields: [
{
label: "IdP Metadata URL",
name: "saml_idp_metadata_url",
required: false,
placeholder: "https://idp.example.com/metadata (use this or the metadata XML below)",
},
{
label: "IdP Metadata XML",
name: "saml_idp_metadata_xml",
required: false,
type: "textarea",
placeholder: "Paste the IdP metadata XML here if you do not have a metadata URL",
},
{
label: "SP Entity ID",
name: "saml_sp_entity_id",
required: false,
placeholder: "Defaults to <proxy base url>/sso/saml/metadata",
},
{
label: "Allow IdP-initiated (unsolicited) responses",
name: "saml_allow_unsolicited",
required: false,
type: "checkbox",
},
],
},
};
// Helper function to render provider fields
@ -97,16 +133,31 @@ export const renderProviderFields = (provider: string) => {
const config = ssoProviderConfigs[provider];
if (!config) return null;
return config.fields.map((field) => (
<Form.Item
key={field.name}
label={field.label}
name={field.name}
rules={[{ required: field.required !== false, message: `Please enter the ${field.label.toLowerCase()}` }]}
>
{field.name.includes("client") ? <Input.Password /> : <TextInput placeholder={field.placeholder} />}
</Form.Item>
));
return config.fields.map((field) => {
const isRequired = field.required !== false;
const rules = isRequired ? [{ required: true, message: `Please enter the ${field.label.toLowerCase()}` }] : [];
let control: React.ReactNode;
if (field.type === "checkbox") {
control = <Checkbox />;
} else if (field.type === "textarea") {
control = <Input.TextArea rows={4} placeholder={field.placeholder} />;
} else if (field.type === "password" || field.name.includes("client")) {
control = <Input.Password />;
} else {
control = <TextInput placeholder={field.placeholder} />;
}
return (
<Form.Item
key={field.name}
label={field.label}
name={field.name}
rules={rules}
valuePropName={field.type === "checkbox" ? "checked" : undefined}
>
{control}
</Form.Item>
);
});
};
const BaseSSOSettingsForm: React.FC<BaseSSOSettingsFormProps> = ({ form, onFormSubmit }) => {

View file

@ -29,6 +29,10 @@ const DeleteSSOSettingsModal: React.FC<DeleteSSOSettingsModalProps> = ({ isVisib
generic_authorization_endpoint: null,
generic_token_endpoint: null,
generic_userinfo_endpoint: null,
saml_idp_metadata_url: null,
saml_idp_metadata_xml: null,
saml_sp_entity_id: null,
saml_allow_unsolicited: null,
proxy_base_url: null,
user_email: null,
sso_provider: null,

View file

@ -181,7 +181,8 @@ vi.mock("@/components/shared/errorUtils", () => ({
parseErrorMessage: vi.fn(),
}));
vi.mock("../utils", () => ({
vi.mock("../utils", async (importOriginal) => ({
...(await importOriginal<typeof import("../utils")>()),
processSSOSettingsPayload: vi.fn(),
}));

View file

@ -5,7 +5,7 @@ import React, { useEffect } from "react";
import BaseSSOSettingsForm from "./BaseSSOSettingsForm";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { parseErrorMessage } from "@/components/shared/errorUtils";
import { processSSOSettingsPayload } from "../utils";
import { detectSSOProvider, processSSOSettingsPayload } from "../utils";
import { useSSOSettings } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
import { useEditSSOSettings } from "@/app/(dashboard)/hooks/sso/useEditSSOSettings";
@ -26,22 +26,7 @@ const EditSSOSettingsModal: React.FC<EditSSOSettingsModalProps> = ({ isVisible,
const ssoData = ssoSettings.data;
// Determine which SSO provider is configured
let selectedProvider = null;
if (ssoData.values.google_client_id) {
selectedProvider = "google";
} else if (ssoData.values.microsoft_client_id) {
selectedProvider = "microsoft";
} else if (ssoData.values.generic_client_id) {
// Check if it looks like Okta based on endpoints
if (
ssoData.values.generic_authorization_endpoint?.includes("okta") ||
ssoData.values.generic_authorization_endpoint?.includes("auth0")
) {
selectedProvider = "okta";
} else {
selectedProvider = "generic";
}
}
const selectedProvider = detectSSOProvider(ssoData.values);
// Extract role mappings if they exist
let roleMappingFields = {};
@ -81,6 +66,9 @@ const EditSSOSettingsModal: React.FC<EditSSOSettingsModalProps> = ({ isVisible,
...ssoData.values,
...roleMappingFields,
...teamMappingFields,
...(ssoData.values.saml_allow_unsolicited != null
? { saml_allow_unsolicited: ssoData.values.saml_allow_unsolicited === "true" }
: {}),
};
// Clear form first, then set values with a small delay to ensure proper initialization

View file

@ -48,6 +48,28 @@ const googleConfiguredValues = {
team_mappings: null,
};
const samlConfiguredValues = {
google_client_id: null,
google_client_secret: null,
microsoft_client_id: null,
microsoft_client_secret: null,
microsoft_tenant: null,
generic_client_id: null,
generic_client_secret: null,
generic_authorization_endpoint: null,
generic_token_endpoint: null,
generic_userinfo_endpoint: null,
proxy_base_url: "https://proxy.example.com",
user_email: null,
ui_access_mode: null,
role_mappings: null,
team_mappings: null,
saml_idp_metadata_url: null,
saml_idp_metadata_xml: "<EntityDescriptor/>",
saml_sp_entity_id: "https://proxy.example.com/sso/saml/metadata",
saml_allow_unsolicited: "true",
};
describe("SSOSettings", () => {
beforeEach(() => {
vi.clearAllMocks();
@ -77,4 +99,20 @@ describe("SSOSettings", () => {
const logo = screen.getByAltText("Google SSO logo");
expect(logo).toHaveAttribute("src", expect.stringContaining("google.svg"));
});
it("renders a SAML configuration as configured instead of the empty placeholder", () => {
mockUseSSOSettings.mockReturnValue({
data: { values: samlConfiguredValues },
isLoading: false,
refetch: vi.fn(),
});
renderSSOSettings();
expect(screen.queryByText("No SSO Configuration Found")).not.toBeInTheDocument();
expect(screen.getByRole("button", { name: /Edit SSO Settings/i })).toBeInTheDocument();
expect(screen.getByText("SAML SSO")).toBeInTheDocument();
expect(screen.getByText("https://proxy.example.com/sso/saml/metadata")).toBeInTheDocument();
expect(screen.getByText("Enabled")).toBeInTheDocument();
});
});

View file

@ -22,10 +22,13 @@ export default function SSOSettings() {
const [isDeleteModalVisible, setIsDeleteModalVisible] = useState(false);
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
const isSSOConfigured =
Boolean(ssoSettings?.values.google_client_id) ||
Boolean(ssoSettings?.values.microsoft_client_id) ||
Boolean(ssoSettings?.values.generic_client_id);
const isSSOConfigured = [
ssoSettings?.values.google_client_id,
ssoSettings?.values.microsoft_client_id,
ssoSettings?.values.generic_client_id,
ssoSettings?.values.saml_idp_metadata_url,
ssoSettings?.values.saml_idp_metadata_xml,
].some(Boolean);
const selectedProvider = ssoSettings?.values ? detectSSOProvider(ssoSettings.values) : null;
const isRoleMappingsEnabled = Boolean(ssoSettings?.values.role_mappings);
@ -154,6 +157,37 @@ export default function SSOSettings() {
: null,
],
},
saml: {
providerText: ssoProviderDisplayNames.saml,
fields: [
{
label: "IdP Metadata URL",
render: (values: SSOSettingsValues) => renderEndpointValue(values.saml_idp_metadata_url),
},
{
label: "IdP Metadata XML",
render: (values: SSOSettingsValues) =>
values.saml_idp_metadata_xml ? (
<Tag>Provided</Tag>
) : (
<span className="text-gray-400 italic">Not configured</span>
),
},
{
label: "SP Entity ID",
render: (values: SSOSettingsValues) => renderEndpointValue(values.saml_sp_entity_id),
},
{
label: "Allow IdP-initiated (unsolicited) responses",
render: (values: SSOSettingsValues) => (
<Tag color={values.saml_allow_unsolicited === "true" ? "green" : "default"}>
{values.saml_allow_unsolicited === "true" ? "Enabled" : "Disabled"}
</Tag>
),
},
{ label: "Proxy Base URL", render: (values: SSOSettingsValues) => renderSimpleValue(values.proxy_base_url) },
],
},
};
const renderSSOSettings = () => {

View file

@ -7,6 +7,7 @@ export const ssoProviderLogoMap: Record<string, string> = {
microsoft: microsoftAzureLogo.src,
okta: "https://www.okta.com/sites/default/files/Okta_Logo_BrightBlue_Medium.png",
generic: "",
saml: "",
};
// SSO Provider display names (consistent between select dropdown and table)
@ -15,6 +16,7 @@ export const ssoProviderDisplayNames: Record<string, string> = {
microsoft: "Microsoft SSO",
okta: "Okta / Auth0 SSO",
generic: "Generic SSO",
saml: "SAML SSO",
};
export const defaultRoleDisplayNames: Record<string, string> = {

View file

@ -1,5 +1,6 @@
import { processSSOSettingsPayload } from "./utils";
import { detectSSOProvider, processSSOSettingsPayload } from "./utils";
import { describe, it, expect } from "vitest";
import type { SSOSettingsValues } from "@/app/(dashboard)/hooks/sso/useSSOSettings";
describe("processSSOSettingsPayload", () => {
describe("without role mappings", () => {
@ -428,3 +429,26 @@ describe("processSSOSettingsPayload", () => {
});
});
});
describe("detectSSOProvider with SAML", () => {
it("returns saml when a SAML IdP metadata URL is configured", () => {
expect(detectSSOProvider({ saml_idp_metadata_url: "https://idp.example.com/metadata" } as SSOSettingsValues)).toBe(
"saml",
);
});
it("returns saml when only inline SAML metadata XML is configured", () => {
expect(detectSSOProvider({ saml_idp_metadata_xml: "<EntityDescriptor/>" } as SSOSettingsValues)).toBe("saml");
});
});
describe("processSSOSettingsPayload with SAML", () => {
it("maps the boolean allow-unsolicited toggle to a 'true'/'false' string", () => {
expect(
processSSOSettingsPayload({ sso_provider: "saml", saml_allow_unsolicited: true }).saml_allow_unsolicited,
).toBe("true");
expect(
processSSOSettingsPayload({ sso_provider: "saml", saml_allow_unsolicited: false }).saml_allow_unsolicited,
).toBe("false");
});
});

View file

@ -22,6 +22,10 @@ export const processSSOSettingsPayload = (formValues: Record<string, any>): Reco
...rest,
};
if (typeof payload.saml_allow_unsolicited === "boolean") {
payload.saml_allow_unsolicited = payload.saml_allow_unsolicited ? "true" : "false";
}
// Add role mappings only if use_role_mappings is checked AND provider supports role mappings
const provider = rest.sso_provider;
const supportsRoleMappings = provider === "okta" || provider === "generic";
@ -81,5 +85,6 @@ export const detectSSOProvider = (values: SSOSettingsValues): string | null => {
}
return "generic";
}
if (values.saml_idp_metadata_url || values.saml_idp_metadata_xml) return "saml";
return null;
};

View file

@ -1,14 +1,12 @@
import type { components } from "@/lib/http/schema";
export interface AgentAttachedKey {
token: string;
key_alias?: string | null;
key_name?: string | null;
}
export interface AgentObjectPermission {
mcp_servers?: string[];
mcp_access_groups?: string[];
mcp_tool_permissions?: Record<string, string[]>;
}
export type AgentObjectPermission = components["schemas"]["AgentObjectPermission"];
export interface Agent {
agent_id: string;

View file

@ -1,357 +1,214 @@
import React, { useState } from "react";
// eslint-disable-next-line no-restricted-imports -- exercising KeyLifecycleSettings requires hosting it in a real antd Form (the component it's built on)
import { Form } from "antd";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi, beforeEach } from "vitest";
import { renderWithProviders, screen } from "../../../tests/test-utils";
import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils";
import KeyLifecycleSettings from "./KeyLifecycleSettings";
vi.mock("antd", () => {
const Option = ({ children, value }: any) => <option value={value}>{children}</option>;
const Select = ({ children, value, onChange, placeholder }: any) => (
<select
data-testid="select"
value={value}
onChange={(e) => onChange(e.target.value)}
data-placeholder={placeholder}
>
{children}
</select>
);
Select.Option = Option;
return {
Select,
Tooltip: ({ children, title }: any) => (
<div data-testid="tooltip" title={title}>
{children}
</div>
),
Switch: ({ checked, onChange }: any) => (
<input type="checkbox" data-testid="switch" checked={checked} onChange={(e) => onChange(e.target.checked)} />
),
Divider: () => <hr data-testid="divider" />,
};
});
const CREATE_PLACEHOLDER = "e.g., 30d or leave empty to never expire";
const EDIT_PLACEHOLDER = "e.g., 30d";
vi.mock("@ant-design/icons", () => ({
InfoCircleOutlined: () => <span data-testid="info-icon">ℹ</span>,
}));
interface HarnessProps {
isCreateMode?: boolean;
onFinish?: (values: Record<string, unknown>) => void;
}
vi.mock("@tremor/react", () => ({
TextInput: ({ value, onValueChange, onChange, placeholder, name, className }: any) => {
const handleChange = (e: React.ChangeEvent<HTMLInputElement>) => {
if (onChange) {
onChange(e);
}
if (onValueChange) {
onValueChange(e.target.value);
}
};
return (
<input
data-testid={name === "duration" ? "duration-input" : "custom-interval-input"}
value={value}
onChange={handleChange}
placeholder={placeholder}
className={className}
const Harness: React.FC<HarnessProps> = ({ isCreateMode = true, onFinish = () => {} }) => {
const [form] = Form.useForm();
const [autoRotationEnabled, setAutoRotationEnabled] = useState(false);
const [rotationInterval, setRotationInterval] = useState("");
const [neverExpire, setNeverExpire] = useState(false);
return (
<Form form={form} onFinish={onFinish}>
<KeyLifecycleSettings
form={form}
autoRotationEnabled={autoRotationEnabled}
onAutoRotationChange={setAutoRotationEnabled}
rotationInterval={rotationInterval}
onRotationIntervalChange={setRotationInterval}
isCreateMode={isCreateMode}
neverExpire={neverExpire}
onNeverExpireChange={setNeverExpire}
/>
);
},
}));
<button type="submit">submit</button>
<button type="button" onClick={() => form.resetFields()}>
reset
</button>
<span data-testid="rotation-interval-value">{rotationInterval}</span>
</Form>
);
};
const getDurationInput = (isCreateMode = true) =>
screen.getByPlaceholderText(isCreateMode ? CREATE_PLACEHOLDER : EDIT_PLACEHOLDER) as HTMLInputElement;
describe("KeyLifecycleSettings", () => {
const mockForm = {
getFieldValue: vi.fn(),
setFieldValue: vi.fn(),
setFieldsValue: vi.fn(),
};
const defaultProps = {
form: mockForm,
autoRotationEnabled: false,
onAutoRotationChange: vi.fn(),
rotationInterval: "",
onRotationIntervalChange: vi.fn(),
isCreateMode: false,
};
beforeEach(() => {
vi.clearAllMocks();
mockForm.getFieldValue.mockReturnValue("");
});
it("should render without crashing", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
it("renders the expiry and auto-rotation sections", () => {
renderWithProviders(<Harness />);
expect(screen.getByText("Key Expiry Settings")).toBeInTheDocument();
expect(screen.getByText("Auto-Rotation Settings")).toBeInTheDocument();
expect(getDurationInput()).toBeInTheDocument();
});
describe("Key Expiry Settings", () => {
it("should render expiry input field", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
it("uses the create-mode placeholder in create mode", () => {
renderWithProviders(<Harness isCreateMode={true} />);
expect(screen.getByPlaceholderText(CREATE_PLACEHOLDER)).toBeInTheDocument();
});
expect(screen.getByText("Expire Key")).toBeInTheDocument();
expect(screen.getByTestId("duration-input")).toBeInTheDocument();
});
it("uses the edit-mode placeholder in edit mode", () => {
renderWithProviders(<Harness isCreateMode={false} />);
expect(screen.getByPlaceholderText(EDIT_PLACEHOLDER)).toBeInTheDocument();
});
it("should show correct placeholder in create mode", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={true} />);
const input = screen.getByTestId("duration-input");
expect(input).toHaveAttribute("placeholder", "e.g., 30d or leave empty to never expire");
});
it("should show correct placeholder in edit mode", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={false} />);
const input = screen.getByTestId("duration-input");
expect(input).toHaveAttribute("placeholder", "e.g., 30d");
});
it("should show correct tooltip in create mode", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={true} />);
const tooltips = screen.getAllByTestId("tooltip");
const expiryTooltip = tooltips.find((tooltip) =>
tooltip.getAttribute("title")?.includes("Leave empty to keep the current expiry unchanged"),
);
expect(expiryTooltip).toBeInTheDocument();
expect(expiryTooltip).toHaveAttribute(
"title",
"Set when this key should expire. Format: 30s (seconds), 30m (minutes), 30h (hours), 30d (days). Leave empty to keep the current expiry unchanged.",
);
});
it("should show correct tooltip in edit mode", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} isCreateMode={false} />);
const tooltips = screen.getAllByTestId("tooltip");
const expiryTooltip = tooltips.find((tooltip) =>
tooltip.getAttribute("title")?.includes("Leave empty to keep the current expiry unchanged"),
);
expect(expiryTooltip).toBeInTheDocument();
expect(expiryTooltip).toHaveAttribute(
"title",
"Set when this key should expire. Format: 30s (seconds), 30m (minutes), 30h (hours), 30d (days). Leave empty to keep the current expiry unchanged.",
);
});
it("should initialize with form value if present", () => {
mockForm.getFieldValue.mockReturnValue("30d");
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
const input = screen.getByTestId("duration-input") as HTMLInputElement;
expect(input.value).toBe("30d");
});
it("should update form using setFieldValue when duration changes", async () => {
describe("duration is a single source of truth (regression for pre-filled value dropped on submit)", () => {
it("submits the duration the user typed", async () => {
const user = userEvent.setup();
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
const onFinish = vi.fn();
renderWithProviders(<Harness onFinish={onFinish} />);
const input = screen.getByTestId("duration-input");
await user.type(input, "60d");
await user.type(getDurationInput(), "1d");
await user.click(screen.getByRole("button", { name: "submit" }));
expect(mockForm.setFieldValue).toHaveBeenCalledWith("duration", "60d");
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(1));
expect(onFinish.mock.calls[0][0]).toMatchObject({ duration: "1d" });
});
it("should update form using setFieldsValue when setFieldValue is not available", async () => {
it("clears the displayed value when the form is reset, so no stale value lingers", async () => {
const user = userEvent.setup();
const formWithoutSetFieldValue = {
getFieldValue: vi.fn().mockReturnValue(""),
setFieldsValue: vi.fn(),
};
renderWithProviders(<KeyLifecycleSettings {...defaultProps} form={formWithoutSetFieldValue} />);
renderWithProviders(<Harness />);
const input = screen.getByTestId("duration-input");
await user.type(input, "90d");
await user.type(getDurationInput(), "1d");
expect(getDurationInput().value).toBe("1d");
expect(formWithoutSetFieldValue.setFieldsValue).toHaveBeenCalledWith({ duration: "90d" });
await user.click(screen.getByRole("button", { name: "reset" }));
await waitFor(() => expect(getDurationInput().value).toBe(""));
});
it("never submits a value that differs from what is displayed after a reset", async () => {
const user = userEvent.setup();
const onFinish = vi.fn();
renderWithProviders(<Harness onFinish={onFinish} />);
// First create: type "1d" and submit -> "1d" is sent.
await user.type(getDurationInput(), "1d");
await user.click(screen.getByRole("button", { name: "submit" }));
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(1));
expect(onFinish.mock.calls[0][0]).toMatchObject({ duration: "1d" });
// Second create: form resets, so the field must show empty AND submit empty.
// The old bug showed a stale "1d" while submitting null/empty.
await user.click(screen.getByRole("button", { name: "reset" }));
await waitFor(() => expect(getDurationInput().value).toBe(""));
await user.click(screen.getByRole("button", { name: "submit" }));
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(2));
expect(onFinish.mock.calls[1][0].duration).not.toBe("1d");
expect(getDurationInput().value).toBe(onFinish.mock.calls[1][0].duration ?? "");
});
});
describe("Auto-Rotation Settings", () => {
it("should render auto-rotation switch", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} />);
expect(screen.getByText("Enable Auto-Rotation")).toBeInTheDocument();
expect(screen.getByTestId("switch")).toBeInTheDocument();
});
it("should show switch as unchecked when autoRotationEnabled is false", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={false} />);
const switchElement = screen.getByTestId("switch") as HTMLInputElement;
expect(switchElement.checked).toBe(false);
});
it("should show switch as checked when autoRotationEnabled is true", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} />);
const switchElement = screen.getByTestId("switch") as HTMLInputElement;
expect(switchElement.checked).toBe(true);
});
it("should call onAutoRotationChange when switch is toggled", async () => {
describe("Never Expire", () => {
it("clears and disables the duration input, then submits an empty duration", async () => {
const user = userEvent.setup();
const onAutoRotationChange = vi.fn();
renderWithProviders(<KeyLifecycleSettings {...defaultProps} onAutoRotationChange={onAutoRotationChange} />);
const onFinish = vi.fn();
renderWithProviders(<Harness isCreateMode={false} onFinish={onFinish} />);
const switchElement = screen.getByTestId("switch");
await user.click(switchElement);
await user.type(getDurationInput(false), "30d");
expect(getDurationInput(false).value).toBe("30d");
expect(onAutoRotationChange).toHaveBeenCalledWith(true);
await user.click(screen.getByRole("checkbox", { name: /never expire/i }));
await waitFor(() => expect(getDurationInput(false).value).toBe(""));
expect(getDurationInput(false)).toBeDisabled();
await user.click(screen.getByRole("button", { name: "submit" }));
await waitFor(() => expect(onFinish).toHaveBeenCalledTimes(1));
expect(onFinish.mock.calls[0][0]).toMatchObject({ duration: "" });
});
});
it("should not show rotation interval section when auto-rotation is disabled", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={false} />);
describe("Auto-Rotation", () => {
it("reveals the rotation interval controls when enabled", async () => {
const user = userEvent.setup();
renderWithProviders(<Harness />);
expect(screen.queryByText("Rotation Interval")).not.toBeInTheDocument();
expect(screen.queryByTestId("select")).not.toBeInTheDocument();
await user.click(screen.getByRole("switch"));
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
});
it("should show rotation interval section when auto-rotation is enabled", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="30d" />);
expect(screen.getByText("Rotation Interval")).toBeInTheDocument();
expect(screen.getByTestId("select")).toBeInTheDocument();
});
it("should show all predefined interval options", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="30d" />);
expect(screen.getByText("7 days")).toBeInTheDocument();
expect(screen.getByText("30 days")).toBeInTheDocument();
expect(screen.getByText("90 days")).toBeInTheDocument();
expect(screen.getByText("180 days")).toBeInTheDocument();
expect(screen.getByText("365 days")).toBeInTheDocument();
expect(screen.getByText("Custom interval")).toBeInTheDocument();
});
it("should display current rotation interval in select", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="90d" />);
const select = screen.getByTestId("select") as HTMLSelectElement;
expect(select.value).toBe("90d");
});
it("should call onRotationIntervalChange when predefined interval is selected", async () => {
it("propagates a selected predefined interval", async () => {
const user = userEvent.setup();
const onRotationIntervalChange = vi.fn();
renderWithProviders(
<KeyLifecycleSettings
{...defaultProps}
autoRotationEnabled={true}
rotationInterval="7d"
onRotationIntervalChange={onRotationIntervalChange}
/>,
);
renderWithProviders(<Harness />);
const select = screen.getByTestId("select");
await user.selectOptions(select, "30d");
await user.click(screen.getByRole("switch"));
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
expect(onRotationIntervalChange).toHaveBeenCalledWith("30d");
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("90 days"));
await waitFor(() => expect(document.querySelector(".ant-select-selection-item")?.textContent).toBe("90 days"));
expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("90d");
});
it("should show custom input when custom option is selected", async () => {
it("shows the custom interval input when Custom interval is selected, without propagating yet", async () => {
const user = userEvent.setup();
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="30d" />);
renderWithProviders(<Harness />);
const select = screen.getByTestId("select");
await user.selectOptions(select, "custom");
await user.click(screen.getByRole("switch"));
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
expect(screen.getByTestId("custom-interval-input")).toBeInTheDocument();
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("Custom interval"));
expect(await screen.findByPlaceholderText("e.g., 1s, 5m, 2h, 14d")).toBeInTheDocument();
expect(screen.getByText("Supported formats: seconds (s), minutes (m), hours (h), days (d)")).toBeInTheDocument();
expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("");
});
it("should hide custom input when predefined interval is selected after custom", async () => {
it("propagates a typed custom interval to the parent", async () => {
const user = userEvent.setup();
const onRotationIntervalChange = vi.fn();
renderWithProviders(
<KeyLifecycleSettings
{...defaultProps}
autoRotationEnabled={true}
rotationInterval="custom-value"
onRotationIntervalChange={onRotationIntervalChange}
/>,
);
renderWithProviders(<Harness />);
const select = screen.getByTestId("select");
await user.selectOptions(select, "7d");
await user.click(screen.getByRole("switch"));
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
expect(screen.queryByTestId("custom-interval-input")).not.toBeInTheDocument();
expect(onRotationIntervalChange).toHaveBeenCalledWith("7d");
});
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("Custom interval"));
it("should call onRotationIntervalChange when custom interval is entered", async () => {
const user = userEvent.setup();
const onRotationIntervalChange = vi.fn();
renderWithProviders(
<KeyLifecycleSettings
{...defaultProps}
autoRotationEnabled={true}
rotationInterval=""
onRotationIntervalChange={onRotationIntervalChange}
/>,
);
const select = screen.getByTestId("select");
await user.selectOptions(select, "custom");
const customInput = screen.getByTestId("custom-interval-input");
const customInput = await screen.findByPlaceholderText("e.g., 1s, 5m, 2h, 14d");
await user.type(customInput, "14d");
expect(onRotationIntervalChange).toHaveBeenCalledWith("14d");
await waitFor(() => expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("14d"));
expect((customInput as HTMLInputElement).value).toBe("14d");
});
it("should show info message when auto-rotation is enabled", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} />);
expect(
screen.getByText(
"When rotation occurs, you'll receive a notification with the new key. The old key will be deactivated after a brief grace period.",
),
).toBeInTheDocument();
});
it("should not show info message when auto-rotation is disabled", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={false} />);
expect(
screen.queryByText(
"When rotation occurs, you'll receive a notification with the new key. The old key will be deactivated after a brief grace period.",
),
).not.toBeInTheDocument();
});
it("should initialize with custom interval input visible when custom interval is provided", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="14d" />);
expect(screen.getByTestId("custom-interval-input")).toBeInTheDocument();
const customInput = screen.getByTestId("custom-interval-input") as HTMLInputElement;
expect(customInput.value).toBe("14d");
});
it("should show custom option selected when custom interval is provided", () => {
renderWithProviders(<KeyLifecycleSettings {...defaultProps} autoRotationEnabled={true} rotationInterval="14d" />);
const select = screen.getByTestId("select") as HTMLSelectElement;
expect(select.value).toBe("custom");
});
it("should not call onRotationIntervalChange when selecting custom option", async () => {
it("hides the custom input and propagates the value when switching back to a predefined interval", async () => {
const user = userEvent.setup();
const onRotationIntervalChange = vi.fn();
renderWithProviders(
<KeyLifecycleSettings
{...defaultProps}
autoRotationEnabled={true}
rotationInterval="30d"
onRotationIntervalChange={onRotationIntervalChange}
/>,
);
renderWithProviders(<Harness />);
const select = screen.getByTestId("select");
await user.selectOptions(select, "custom");
await user.click(screen.getByRole("switch"));
await waitFor(() => expect(screen.getByText("Rotation Interval")).toBeInTheDocument());
expect(onRotationIntervalChange).not.toHaveBeenCalled();
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("Custom interval"));
const customInput = await screen.findByPlaceholderText("e.g., 1s, 5m, 2h, 14d");
await user.type(customInput, "14d");
await waitFor(() => expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("14d"));
await user.click(screen.getByRole("combobox"));
await user.click(await screen.findByText("7 days"));
await waitFor(() => expect(screen.getByTestId("rotation-interval-value")).toHaveTextContent("7d"));
expect(screen.queryByPlaceholderText("e.g., 1s, 5m, 2h, 14d")).not.toBeInTheDocument();
});
});
});

View file

@ -1,5 +1,5 @@
import React, { useState } from "react";
import { Select, Tooltip, Divider, Switch, Checkbox } from "antd";
import { Select, Tooltip, Divider, Switch, Checkbox, Form } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { TextInput } from "@tremor/react";
@ -34,7 +34,6 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
const [showCustomInput, setShowCustomInput] = useState(isCustomInterval);
const [customInterval, setCustomInterval] = useState(isCustomInterval ? rotationInterval : "");
const [durationValue, setDurationValue] = useState<string>(form?.getFieldValue?.("duration") || "");
const handleIntervalChange = (value: string) => {
if (value === "custom") {
@ -53,14 +52,6 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
onRotationIntervalChange(value);
};
const handleDurationChange = (value: string) => {
setDurationValue(value);
if (form && typeof form.setFieldValue === "function") {
form.setFieldValue("duration", value);
} else if (form && typeof form.setFieldsValue === "function") {
form.setFieldsValue({ duration: value });
}
};
return (
<div className="space-y-6">
{/* Key Expiry Section */}
@ -80,7 +71,6 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
const checked = e.target.checked;
onNeverExpireChange(checked);
if (checked) {
setDurationValue("");
if (form && typeof form.setFieldValue === "function") {
form.setFieldValue("duration", "");
} else if (form && typeof form.setFieldsValue === "function") {
@ -94,14 +84,13 @@ const KeyLifecycleSettings: React.FC<KeyLifecycleSettingsProps> = ({
</Checkbox>
)}
</label>
<TextInput
name="duration"
placeholder={isCreateMode ? "e.g., 30d or leave empty to never expire" : "e.g., 30d"}
className="w-full"
value={durationValue}
onValueChange={handleDurationChange}
disabled={!isCreateMode && neverExpire}
/>
<Form.Item name="duration" noStyle initialValue="">
<TextInput
placeholder={isCreateMode ? "e.g., 30d or leave empty to never expire" : "e.g., 30d"}
className="w-full"
disabled={!isCreateMode && neverExpire}
/>
</Form.Item>
</div>
</div>

View file

@ -1,6 +1,7 @@
import { Setter } from "@/types";
import { useEffect, useState } from "react";
import { keyListCall, Member, Organization } from "../networking";
import type { ObjectPermission } from "../object_permission_types";
export interface Team {
team_id: string;
@ -90,15 +91,7 @@ export interface KeyResponse {
user_tpm_limit: number;
user_rpm_limit: number;
user_email: string;
object_permission?: {
object_permission_id: string;
mcp_servers: string[];
mcp_access_groups?: string[];
mcp_tool_permissions?: Record<string, string[]>;
vector_stores: string[];
agents?: string[];
agent_access_groups?: string[];
};
object_permission?: ObjectPermission | null;
access_group_ids?: string[];
budget_fallbacks?: Record<string, string[]>;
budget_limits?: Array<{ budget_duration: string; max_budget: number; reset_at?: string }>;

View file

@ -28,6 +28,7 @@ import { TagNewRequest, TagUpdateRequest, TagListResponse, TagInfoResponse } fro
import { Team } from "./key_team_helpers/key_list";
import { EmailEventSettingsResponse, EmailEventSettingsUpdateRequest } from "./email_events/types";
import type { SkillRegisterRequest } from "./claude_code_plugins/types";
import type { ObjectPermission } from "./object_permission_types";
import { jsonFields } from "./common_components/check_openapi_schema";
import NotificationsManager from "./molecules/notifications_manager";
import type { MCPUserEnvVarsStatus } from "./mcp_tools/types";
@ -208,13 +209,7 @@ export interface Organization {
teams: any[] | null;
users: any[] | null;
members: any[] | null;
object_permission?: {
object_permission_id: string;
mcp_servers: string[];
mcp_access_groups?: string[];
mcp_toolsets?: string[];
vector_stores: string[];
};
object_permission?: ObjectPermission | null;
}
export interface CredentialItem {
@ -1184,35 +1179,6 @@ export const organizationInfoCall = async (accessToken: string, organizationID:
}
};
export const organizationCreateCall = async (
accessToken: string,
formValues: Record<string, any>, // Assuming formValues is an object
) => {
try {
if (formValues.metadata) {
// if there's an exception JSON.parse, show it in the message
try {
formValues.metadata = JSON.parse(formValues.metadata);
} catch (error) {
console.error("Failed to parse metadata:", error);
throw new Error("Failed to parse metadata: " + error);
}
}
const data = await apiClient.post(`/organization/new`, {
accessToken,
body: {
...formValues, // Include formValues in the request body
},
});
return data;
// Handle success - you might want to update some state or UI based on the created key
} catch (error) {
console.error("Failed to create key:", error);
throw error;
}
};
export const organizationUpdateCall = async (
accessToken: string,
formValues: Record<string, any>, // Assuming formValues is an object

View file

@ -0,0 +1,3 @@
import type { components } from "@/lib/http/schema";
export type ObjectPermission = Partial<components["schemas"]["LiteLLM_ObjectPermissionTable"]>;

View file

@ -3,21 +3,10 @@ import { Text } from "@tremor/react";
import VectorStorePermissions from "./permissions/VectorStorePermissions";
import MCPServerPermissions from "./permissions/MCPServerPermissions";
import AgentPermissions from "./permissions/AgentPermissions";
interface ObjectPermission {
object_permission_id: string;
mcp_servers: string[];
mcp_access_groups?: string[];
mcp_tool_permissions?: Record<string, string[]>;
mcp_toolsets?: string[];
vector_stores: string[];
agents?: string[];
agent_access_groups?: string[];
search_tools?: string[];
}
import type { ObjectPermission } from "./object_permission_types";
interface ObjectPermissionsViewProps {
objectPermission?: ObjectPermission;
objectPermission?: ObjectPermission | null;
variant?: "card" | "inline";
className?: string;
accessToken?: string | null;

View file

@ -443,6 +443,36 @@ describe("CreateKey", () => {
});
});
it("should include mcp_toolsets in keyCreateCall payload when only toolsets are selected", async () => {
renderWithProviders(<CreateKey {...defaultProps} />);
act(() => {
fireEvent.click(screen.getByRole("button", { name: /create new key/i }));
});
await waitFor(() => {
expect(screen.getByRole("button", { name: /create key/i })).toBeInTheDocument();
});
act(() => {
formMock.setFieldValue("key_alias", "Test Key");
formMock.setFieldValue("allowed_mcp_servers_and_groups", {
servers: [],
accessGroups: [],
toolsets: ["ts-1"],
});
});
act(() => {
fireEvent.click(screen.getByRole("button", { name: /create key/i }));
});
await waitFor(() => {
expect(mockKeyCreateCall).toHaveBeenCalled();
});
expect(mockKeyCreateCall.mock.calls[0][2].object_permission?.mcp_toolsets).toEqual(["ts-1"]);
});
it("should prefill models when provided without team_id", async () => {
renderWithProviders(
<CreateKey

View file

@ -468,18 +468,22 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
if (
formValues.allowed_mcp_servers_and_groups &&
(formValues.allowed_mcp_servers_and_groups.servers?.length > 0 ||
formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0)
formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0 ||
formValues.allowed_mcp_servers_and_groups.toolsets?.length > 0)
) {
if (!formValues.object_permission) {
formValues.object_permission = {};
}
const { servers, accessGroups } = formValues.allowed_mcp_servers_and_groups;
const { servers, accessGroups, toolsets } = formValues.allowed_mcp_servers_and_groups;
if (servers && servers.length > 0) {
formValues.object_permission.mcp_servers = servers;
}
if (accessGroups && accessGroups.length > 0) {
formValues.object_permission.mcp_access_groups = accessGroups;
}
if (toolsets && toolsets.length > 0) {
formValues.object_permission.mcp_toolsets = toolsets;
}
// Remove the original field as it's now part of object_permission
delete formValues.allowed_mcp_servers_and_groups;
}
@ -1640,9 +1644,6 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
/>
</div>
</AccordionBody>
<Form.Item name="duration" hidden initialValue={null}>
<Input />
</Form.Item>
</Accordion>
<Accordion className="mt-4 mb-4">
<AccordionHeader>

View file

@ -0,0 +1,203 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import React from "react";
import { describe, expect, it, vi } from "vitest";
vi.mock("@/components/molecules/notifications_manager", () => ({
__esModule: true,
default: { success: vi.fn(), fromBackend: vi.fn() },
}));
vi.mock("@/components/ModelSelect/ModelSelect", () => ({
ModelSelect: ({ onChange }: { onChange: (values: string[]) => void }) => (
<button type="button" onClick={() => onChange(["gpt-5.2"])}>
set-models
</button>
),
}));
vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({
__esModule: true,
default: ({ onChange }: { onChange: (values: string[]) => void }) => (
<button type="button" onClick={() => onChange(["vs-1"])}>
set-vector-stores
</button>
),
}));
vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({
__esModule: true,
default: ({
onChange,
}: {
onChange: (values: { servers: string[]; accessGroups: string[]; toolsets: string[] }) => void;
}) => (
<button type="button" onClick={() => onChange({ servers: ["srv-1"], accessGroups: [], toolsets: ["ts-1"] })}>
set-mcp
</button>
),
}));
import { OrgCreateDialog } from "./OrgCreateDialog";
const Harness = ({ createOrganization }: { createOrganization: (body: unknown) => Promise<unknown> }) => {
const [open, setOpen] = React.useState(true);
return (
<>
<button type="button" onClick={() => setOpen(true)}>
reopen
</button>
<OrgCreateDialog open={open} onOpenChange={setOpen} accessToken="token" createOrganization={createOrganization} />
</>
);
};
const renderDialog = (overrides?: { createOrganization?: ReturnType<typeof vi.fn> }) => {
const createOrganization = overrides?.createOrganization ?? vi.fn().mockResolvedValue({});
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
render(
<QueryClientProvider client={queryClient}>
<Harness createOrganization={createOrganization} />
</QueryClientProvider>,
);
return { createOrganization };
};
describe("OrgCreateDialog", () => {
it("blocks submit and shows an error when the name is missing", async () => {
const user = userEvent.setup();
const { createOrganization } = renderDialog();
await user.click(screen.getByRole("button", { name: "Create Organization" }));
expect(await screen.findByRole("alert")).toHaveTextContent("Please input an organization name");
expect(createOrganization).not.toHaveBeenCalled();
});
it("sends only alias and models for a minimal create and closes the dialog", async () => {
const user = userEvent.setup();
const { createOrganization } = renderDialog();
await user.type(screen.getByLabelText("Organization Name"), "new-org");
await user.click(screen.getByRole("button", { name: "Create Organization" }));
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
expect(createOrganization.mock.calls[0][0]).toStrictEqual({ organization_alias: "new-org", models: [] });
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
});
it("maps selectors and limits into the create body", async () => {
const user = userEvent.setup();
const { createOrganization } = renderDialog();
await user.type(screen.getByLabelText("Organization Name"), "new-org");
await user.click(screen.getByRole("button", { name: "set-models" }));
await user.type(screen.getByLabelText("Tokens per minute Limit (TPM)"), "1000");
await user.click(screen.getByRole("button", { name: "set-vector-stores" }));
await user.click(screen.getByRole("button", { name: "set-mcp" }));
await user.click(screen.getByRole("button", { name: "Create Organization" }));
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
const expectedBody = {
organization_alias: "new-org",
models: ["gpt-5.2"],
tpm_limit: 1000,
object_permission: {
vector_stores: ["vs-1"],
mcp_servers: ["srv-1"],
mcp_toolsets: ["ts-1"],
},
};
expect(createOrganization.mock.calls[0][0]).toStrictEqual(expectedBody);
});
it("blocks submit and shows an error for invalid metadata JSON", async () => {
const user = userEvent.setup();
const { createOrganization } = renderDialog();
await user.type(screen.getByLabelText("Organization Name"), "new-org");
await user.type(screen.getByLabelText("Metadata"), "not json");
await user.click(screen.getByRole("button", { name: "Create Organization" }));
expect(await screen.findByRole("alert")).toHaveTextContent("Metadata must be a valid JSON object");
expect(createOrganization).not.toHaveBeenCalled();
});
it("keeps the dialog open with the entered values when the create fails", async () => {
const user = userEvent.setup();
const { createOrganization } = renderDialog({
createOrganization: vi.fn().mockRejectedValue(new Error("boom")),
});
await user.type(screen.getByLabelText("Organization Name"), "new-org");
await user.click(screen.getByRole("button", { name: "Create Organization" }));
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
expect(screen.getByLabelText("Organization Name")).toHaveValue("new-org");
});
it("resets the form when the dialog is cancelled and reopened", async () => {
const user = userEvent.setup();
renderDialog();
await user.type(screen.getByLabelText("Organization Name"), "abandoned");
await user.click(screen.getByRole("button", { name: "Cancel" }));
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
await user.click(screen.getByRole("button", { name: "reopen" }));
expect(screen.getByLabelText("Organization Name")).toHaveValue("");
});
it("resets the form when the dialog is dismissed with Escape and reopened", async () => {
const user = userEvent.setup();
renderDialog();
await user.type(screen.getByLabelText("Organization Name"), "abandoned");
await user.keyboard("{Escape}");
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
await user.click(screen.getByRole("button", { name: "reopen" }));
expect(screen.getByLabelText("Organization Name")).toHaveValue("");
});
it("cannot be dismissed while a create is pending, then closes once on success", async () => {
const user = userEvent.setup();
let resolveCreate: (value: unknown) => void = () => {};
const createOrganization = vi.fn().mockImplementation(
() =>
new Promise((resolve) => {
resolveCreate = resolve;
}),
);
renderDialog({ createOrganization });
await user.type(screen.getByLabelText("Organization Name"), "new-org");
await user.keyboard("{Enter}");
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
await user.keyboard("{Escape}");
expect(screen.getByLabelText("Organization Name")).toHaveValue("new-org");
resolveCreate({});
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
});
it("does not fire a second create while one is pending", async () => {
const user = userEvent.setup();
let resolveCreate: (value: unknown) => void = () => {};
const createOrganization = vi.fn().mockImplementation(
() =>
new Promise((resolve) => {
resolveCreate = resolve;
}),
);
renderDialog({ createOrganization });
await user.type(screen.getByLabelText("Organization Name"), "new-org");
await user.keyboard("{Enter}");
await waitFor(() => expect(createOrganization).toHaveBeenCalledTimes(1));
await user.keyboard("{Enter}");
expect(createOrganization).toHaveBeenCalledTimes(1);
resolveCreate({});
await waitFor(() => expect(screen.queryByLabelText("Organization Name")).not.toBeInTheDocument());
});
});

View file

@ -0,0 +1,186 @@
"use client";
import { useMutation, useQueryClient } from "@tanstack/react-query";
import * as React from "react";
import { organizationKeys } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { FieldGroup } from "@/components/shared/form/field";
import { FormField } from "@/components/shared/form/FormField";
import { Button } from "@/components/ui/button";
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
import { Input } from "@/components/ui/input";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Textarea } from "@/components/ui/textarea";
import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector";
import { useZodForm } from "@/lib/forms/useZodForm";
import { fetchClient } from "@/lib/http/api";
import { BUDGET_DURATION_OPTIONS, NO_RESET } from "../org-settings/OrgSettingsForm";
import { orgSettingsSchema } from "../org-settings/schema";
import { buildOrgCreateBody, emptyOrgFormValues, type OrgCreateBody } from "./mapper";
const defaultCreateOrganization = async (body: OrgCreateBody): Promise<unknown> => {
const { data } = await fetchClient.POST("/organization/new", { body });
return data;
};
interface OrgCreateDialogProps {
open: boolean;
onOpenChange: (open: boolean) => void;
accessToken: string;
createOrganization?: (body: OrgCreateBody) => Promise<unknown>;
}
export const OrgCreateDialog = ({
open,
onOpenChange,
accessToken,
createOrganization = defaultCreateOrganization,
}: OrgCreateDialogProps) => {
const queryClient = useQueryClient();
const form = useZodForm(orgSettingsSchema, { defaultValues: emptyOrgFormValues });
const closeAndReset = () => {
form.reset(emptyOrgFormValues);
onOpenChange(false);
};
const mutation = useMutation({
mutationFn: (body: OrgCreateBody) => createOrganization(body),
onSuccess: () => {
NotificationsManager.success("Organization created successfully");
queryClient.invalidateQueries({ queryKey: organizationKeys.all });
closeAndReset();
},
onError: (error: unknown) =>
NotificationsManager.fromBackend(error instanceof Error ? error.message : "Failed to create organization"),
});
const handleOpenChange = (nextOpen: boolean) => {
if (!nextOpen && mutation.isPending) return;
if (!nextOpen) {
form.reset(emptyOrgFormValues);
}
onOpenChange(nextOpen);
};
const onSubmit = form.handleSubmit((values) => {
if (mutation.isPending) return;
mutation.mutate(buildOrgCreateBody(values));
});
return (
<Dialog open={open} onOpenChange={handleOpenChange}>
<DialogContent className="sm:max-w-3xl max-h-[90vh] overflow-y-auto">
<DialogHeader>
<DialogTitle>Create Organization</DialogTitle>
</DialogHeader>
<form onSubmit={onSubmit}>
<FieldGroup>
<FormField control={form.control} name="organization_alias" label="Organization Name">
{({ ref, ...field }) => <Input {...field} ref={ref} />}
</FormField>
<FormField control={form.control} name="models" label="Models">
{(field) => (
<ModelSelect
value={field.value}
onChange={field.onChange}
context="organization"
options={{ includeSpecialOptions: true, showAllProxyModelsOverride: true }}
/>
)}
</FormField>
<FormField control={form.control} name="max_budget" label="Max Budget (USD)">
{({ ref, ...field }) => <Input {...field} ref={ref} type="number" step={0.01} min={0} />}
</FormField>
<FormField control={form.control} name="budget_duration" label="Reset Budget">
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
<Select
items={BUDGET_DURATION_OPTIONS}
value={value === "" ? NO_RESET : value}
onValueChange={(selected) => onChange(selected === NO_RESET ? "" : selected)}
>
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy}>
<SelectValue />
</SelectTrigger>
<SelectContent>
{BUDGET_DURATION_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</FormField>
<FormField control={form.control} name="tpm_limit" label="Tokens per minute Limit (TPM)">
{({ ref, ...field }) => <Input {...field} ref={ref} type="number" step={1} min={0} />}
</FormField>
<FormField control={form.control} name="rpm_limit" label="Requests per minute Limit (RPM)">
{({ ref, ...field }) => <Input {...field} ref={ref} type="number" step={1} min={0} />}
</FormField>
<FormField
control={form.control}
name="vector_stores"
label="Allowed Vector Stores"
description="Select vector stores this organization can access. Leave empty for access to all vector stores"
>
{(field) => (
<VectorStoreSelector
value={field.value}
onChange={field.onChange}
accessToken={accessToken}
placeholder="Select vector stores (optional)"
/>
)}
</FormField>
<FormField
control={form.control}
name="mcp"
label="Allowed MCP Servers"
description="Select MCP servers, access groups, and toolsets this organization can access. Leave empty for access to all"
>
{(field) => (
<MCPServerSelector
value={field.value}
onChange={field.onChange}
accessToken={accessToken}
placeholder="Select MCP servers and access groups (optional)"
/>
)}
</FormField>
<FormField control={form.control} name="metadata" label="Metadata">
{({ ref, ...field }) => <Textarea {...field} ref={ref} rows={4} />}
</FormField>
</FieldGroup>
<DialogFooter className="mt-6">
<Button
type="button"
variant="outline"
onClick={() => handleOpenChange(false)}
disabled={mutation.isPending}
>
Cancel
</Button>
<Button type="submit" disabled={mutation.isPending}>
{mutation.isPending ? "Creating..." : "Create Organization"}
</Button>
</DialogFooter>
</form>
</DialogContent>
</Dialog>
);
};

View file

@ -0,0 +1,59 @@
import { describe, expect, it } from "vitest";
import { buildOrgCreateBody, emptyOrgFormValues } from "./mapper";
describe("buildOrgCreateBody", () => {
it("sends only alias and models for a minimal form", () => {
expect(buildOrgCreateBody({ ...emptyOrgFormValues, organization_alias: "acme" })).toStrictEqual({
organization_alias: "acme",
models: [],
});
});
it("maps every field when the whole form is filled", () => {
const filledForm = {
organization_alias: "acme",
models: ["gpt-5.2"],
max_budget: "12.5",
budget_duration: "30d",
tpm_limit: "1000",
rpm_limit: "50",
vector_stores: ["vs-1"],
mcp: { servers: ["srv-1"], accessGroups: ["ag-1"], toolsets: ["ts-1"] },
metadata: '{"env": "prod"}',
};
const expectedBody = {
organization_alias: "acme",
models: ["gpt-5.2"],
max_budget: 12.5,
budget_duration: "30d",
tpm_limit: 1000,
rpm_limit: 50,
metadata: { env: "prod" },
object_permission: {
vector_stores: ["vs-1"],
mcp_servers: ["srv-1"],
mcp_access_groups: ["ag-1"],
mcp_toolsets: ["ts-1"],
},
};
expect(buildOrgCreateBody(filledForm)).toStrictEqual(expectedBody);
});
it("includes only the non-empty grant lists in object_permission", () => {
expect(
buildOrgCreateBody({
...emptyOrgFormValues,
organization_alias: "acme",
mcp: { servers: [], accessGroups: [], toolsets: ["ts-1"] },
}).object_permission,
).toStrictEqual({ mcp_toolsets: ["ts-1"] });
});
it("parses metadata into an object instead of sending the raw string", () => {
expect(
buildOrgCreateBody({ ...emptyOrgFormValues, organization_alias: "acme", metadata: '{"a": 1}' }).metadata,
).toStrictEqual({ a: 1 });
});
});

View file

@ -0,0 +1,45 @@
import { z } from "zod/v4";
import type { components } from "@/lib/http/schema";
import type { OrgSettingsFormValues } from "../org-settings/schema";
export type OrgCreateBody = components["schemas"]["NewOrganizationRequest"];
export const emptyOrgFormValues: OrgSettingsFormValues = {
organization_alias: "",
models: [],
max_budget: "",
budget_duration: "",
tpm_limit: "",
rpm_limit: "",
vector_stores: [],
mcp: { servers: [], accessGroups: [], toolsets: [] },
metadata: "",
};
const metadataRecordSchema = z.record(z.string(), z.unknown());
const objectPermissionFromValues = (values: OrgSettingsFormValues): OrgCreateBody["object_permission"] => {
const grants = {
...(values.vector_stores.length > 0 && { vector_stores: values.vector_stores }),
...(values.mcp.servers.length > 0 && { mcp_servers: values.mcp.servers }),
...(values.mcp.accessGroups.length > 0 && { mcp_access_groups: values.mcp.accessGroups }),
...(values.mcp.toolsets.length > 0 && { mcp_toolsets: values.mcp.toolsets }),
};
return Object.keys(grants).length > 0 ? grants : undefined;
};
export const buildOrgCreateBody = (values: OrgSettingsFormValues): OrgCreateBody => {
const objectPermission = objectPermissionFromValues(values);
return {
organization_alias: values.organization_alias,
models: values.models,
...(values.max_budget.trim() !== "" && { max_budget: Number(values.max_budget) }),
...(values.tpm_limit.trim() !== "" && { tpm_limit: Number(values.tpm_limit) }),
...(values.rpm_limit.trim() !== "" && { rpm_limit: Number(values.rpm_limit) }),
...(values.budget_duration !== "" && { budget_duration: values.budget_duration }),
...(values.metadata.trim() !== "" && { metadata: metadataRecordSchema.parse(JSON.parse(values.metadata)) }),
...(objectPermission !== undefined && { object_permission: objectPermission }),
};
};

View file

@ -22,9 +22,9 @@ import { fetchClient } from "@/lib/http/api";
import { buildOrgPatch, orgToForm, type OrgPatchBody } from "./mapper";
import { orgSettingsSchema } from "./schema";
const NO_RESET = "never";
export const NO_RESET = "never";
const BUDGET_DURATION_OPTIONS = [
export const BUDGET_DURATION_OPTIONS = [
{ value: NO_RESET, label: "No reset" },
{ value: "24h", label: "daily" },
{ value: "7d", label: "weekly" },
@ -102,6 +102,7 @@ export const OrgSettingsForm = ({
<FormField control={form.control} name="budget_duration" label="Reset Budget">
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
<Select
items={BUDGET_DURATION_OPTIONS}
value={value === "" ? NO_RESET : value}
onValueChange={(selected) => onChange(selected === NO_RESET ? "" : selected)}
>

View file

@ -17,6 +17,7 @@ import {
import { useGuardrails, GuardrailListItem } from "@/app/(dashboard)/hooks/guardrails/useGuardrails";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { mapEmptyStringToNull } from "@/utils/keyUpdateUtils";
import type { ObjectPermission } from "@/components/object_permission_types";
import { isProxyAdminRole } from "@/utils/roles";
import {
EditOutlined,
@ -118,17 +119,7 @@ export interface TeamData {
router_settings?: Record<string, any>;
guardrails?: string[];
policies?: string[];
object_permission?: {
object_permission_id: string;
mcp_servers: string[];
mcp_access_groups?: string[];
mcp_tool_permissions?: Record<string, string[]>;
mcp_toolsets?: string[];
vector_stores: string[];
agents?: string[];
agent_access_groups?: string[];
search_tools?: string[];
};
object_permission?: ObjectPermission | null;
team_member_budget_table: {
max_budget: number;
budget_duration: string;

Some files were not shown because too many files have changed in this diff Show more